Enable CP shared-KV compute padding without inflating cache state
Tiny extend requests can leave most CP lanes without query work, which has been tied to hangs and accept-length regressions. This change introduces a dual valid/compute metadata contract: forward paths may materialize compute-padded rows, while cache, current reuse, direct write, HiCache backup, and load remain valid/page based. The implementation keeps radix/HiCache/device allocation on real page extents, filters dummy compute rows before MLA/index cache writes and current reuse, makes top-k/index consume compute rows while compacting valid rows, and opens tiny CP shared-KV in-seq split through compute padding. The accompanying plan document records the contract and P1-P7 evidence. Constraint: CP shared KV and HiCache must stay page-granular; dummy compute rows must not allocate, write, backup, or load KV cache. Constraint: Avoid silent fallback and avoid adding collectives on hot paths. Rejected: Pad cache allocations to cp_size pages | would waste KV capacity and pollute radix/HiCache state. Rejected: Keep tiny suffixes out of CP split | preserves the zero-lane behavior that compute padding is meant to remove. Confidence: medium Scope-risk: broad Directive: Do not route compute-padded dummy rows into out_cache_loc, current reuse, HiCache reservation, or backup descriptors; keep valid/cache metadata explicit. Tested: Remote g0034 container targeted P7 tests: 3 passed, 3 warnings. Tested: Remote g0034 container full unit slice: PYTHONPATH=python python -m pytest -q test/registered/unit/layers/test_nsa_cp_utils.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py => 214 passed, 5 warnings, 2 subtests passed. Tested: Local py_compile for touched P7 test file. Not-tested: Latest CUDA/ETE traffic validation for dummy top-k rows, accept len, output len, and detokenizer hang behavior.
This commit is contained in:
@@ -12,6 +12,7 @@ from sglang.srt.layers.attention.nsa.utils import (
|
||||
build_batch_page_aligned_in_seq_split_plan,
|
||||
build_page_aligned_cache_extent,
|
||||
_get_in_seq_last_token_owner_and_offset,
|
||||
_build_batch_metadata_from_plan,
|
||||
build_page_aligned_in_seq_split_list,
|
||||
build_token_balanced_in_seq_split_list,
|
||||
can_cp_split,
|
||||
@@ -19,6 +20,7 @@ from sglang.srt.layers.attention.nsa.utils import (
|
||||
cp_collect_last_token_hidden,
|
||||
cp_split_and_rebuild_1d,
|
||||
cp_split_and_rebuild_data,
|
||||
cp_split_and_rebuild_position,
|
||||
_torch_batch_in_seq_all_gather_rerange,
|
||||
get_cp_shared_kv_batch_plan,
|
||||
get_cp_shared_kv_local_out_cache_loc,
|
||||
@@ -26,6 +28,7 @@ from sglang.srt.layers.attention.nsa.utils import (
|
||||
get_cp_local_embedding_padded_token_count,
|
||||
pad_cp_local_input_ids_for_embedding,
|
||||
prepare_input_dp_with_cp_dsa,
|
||||
select_cp_local_valid_rows_for_cache_write,
|
||||
split_tensor_by_cp_batch_plan,
|
||||
split_in_seq_cp_local_pair,
|
||||
)
|
||||
@@ -242,7 +245,7 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
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):
|
||||
def test_can_cp_split_uses_compute_padding_for_short_radix_hit_suffix(self):
|
||||
class Mode:
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
@@ -265,9 +268,11 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
self.assertFalse(can_cp_split(128, 8, True, forward_batch))
|
||||
self.assertTrue(can_cp_split(128, 8, True, forward_batch))
|
||||
|
||||
def test_can_cp_split_skips_cp_when_page_units_do_not_cover_all_lanes(self):
|
||||
def test_can_cp_split_uses_compute_padding_when_page_units_do_not_cover_all_lanes(
|
||||
self,
|
||||
):
|
||||
class Mode:
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
@@ -290,9 +295,9 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
self.assertFalse(can_cp_split(256, 8, True, forward_batch))
|
||||
self.assertTrue(can_cp_split(256, 8, True, forward_batch))
|
||||
|
||||
def test_can_cp_split_skips_current_only_when_page_units_do_not_cover_all_lanes(
|
||||
def test_can_cp_split_uses_compute_padding_for_current_only_one_page_suffix(
|
||||
self,
|
||||
):
|
||||
class Mode:
|
||||
@@ -317,7 +322,34 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
self.assertFalse(can_cp_split(64, 8, True, forward_batch))
|
||||
self.assertTrue(can_cp_split(64, 8, True, forward_batch))
|
||||
|
||||
def test_can_cp_split_uses_compute_padding_per_request_for_batched_tiny_suffix(
|
||||
self,
|
||||
):
|
||||
class Mode:
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
extend_seq_lens_cpu=[65, 64],
|
||||
extend_prefix_lens_cpu=[54464, 8192],
|
||||
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(129, 8, True, forward_batch))
|
||||
|
||||
def test_can_cp_split_fails_on_non_page_aligned_cp_shared_prefix(self):
|
||||
class Mode:
|
||||
@@ -445,6 +477,70 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
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_exposes_compute_padding_without_inflating_valid_cache_extent(
|
||||
self,
|
||||
):
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[40320],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
|
||||
self.assertTrue(plan.compute_padding_enabled)
|
||||
self.assertEqual(
|
||||
plan.request_valid_split_lists,
|
||||
[[64, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]],
|
||||
)
|
||||
self.assertEqual(
|
||||
plan.request_compute_split_lists,
|
||||
[[64, 64, 64, 64, 64, 64, 64, 64, 0, 0, 0, 0, 0, 0, 0, 0]],
|
||||
)
|
||||
self.assertEqual(plan.request_valid_padded_pages, [2])
|
||||
self.assertEqual(plan.request_valid_padded_tokens, [128])
|
||||
self.assertEqual(plan.request_compute_padded_pages, [8])
|
||||
self.assertEqual(plan.request_compute_padded_tokens, [512])
|
||||
self.assertEqual(plan.request_compute_padding_tokens, [447])
|
||||
self.assertEqual(plan.request_compute_rank_local_tokens, [64])
|
||||
self.assertEqual(plan.request_compute_rank_local_offsets, [0])
|
||||
self.assertEqual(plan.request_valid_rank_local_tokens, [1])
|
||||
self.assertEqual(plan.request_valid_rank_local_offsets, [0])
|
||||
self.assertEqual(plan.request_last_token_owner, [1])
|
||||
self.assertEqual(plan.request_last_token_local_offset, [0])
|
||||
|
||||
# Compatibility aliases for cache/page accounting stay valid-token
|
||||
# based. Query-length metadata is split separately below: attention and
|
||||
# top-k consume compute rows, cache/current paths consume valid rows.
|
||||
self.assertEqual(plan.request_split_lists, plan.request_valid_split_lists)
|
||||
self.assertEqual(plan.request_padded_pages, plan.request_valid_padded_pages)
|
||||
self.assertEqual(plan.request_actual_seq_q_prev, [64])
|
||||
self.assertEqual(plan.request_actual_seq_q_next, [0])
|
||||
self.assertEqual(plan.request_valid_seq_q_prev, [1])
|
||||
self.assertEqual(plan.request_valid_seq_q_next, [0])
|
||||
self.assertEqual(plan.request_compute_seq_q_prev, [64])
|
||||
self.assertEqual(plan.request_compute_seq_q_next, [0])
|
||||
|
||||
def test_batch_plan_compute_padding_is_per_request_not_batch_total(self):
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65, 1024],
|
||||
prefix_lens=[40320, 8192],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=0,
|
||||
)
|
||||
|
||||
self.assertTrue(plan.compute_padding_enabled)
|
||||
self.assertEqual(plan.request_valid_padded_pages, [2, 16])
|
||||
self.assertEqual(plan.request_compute_padded_pages, [8, 16])
|
||||
self.assertEqual(plan.request_compute_padded_tokens, [512, 1024])
|
||||
self.assertEqual(plan.request_compute_padding_tokens, [447, 0])
|
||||
self.assertEqual(plan.request_compute_rank_local_tokens, [64, 128])
|
||||
self.assertEqual(plan.request_compute_rank_local_offsets, [0, 64])
|
||||
self.assertEqual(plan.request_valid_rank_local_tokens, [64, 128])
|
||||
self.assertEqual(plan.request_valid_rank_local_offsets, [0, 64])
|
||||
self.assertEqual(plan.request_last_token_owner, [1, 0])
|
||||
|
||||
def test_batch_plan_stable_helpers_split_and_build_page_owner_plan(self):
|
||||
import torch
|
||||
|
||||
@@ -461,12 +557,15 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
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])
|
||||
self.assertEqual(local_1d.tolist(), [0, 0, 0, 0, 8, 9, 10, 11, 12, 0, 0, 0])
|
||||
|
||||
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)))
|
||||
self.assertEqual(
|
||||
local_data[:, 0].tolist(),
|
||||
[0, 0, 0, 0, 16, 18, 20, 22, 24, 0, 0, 0],
|
||||
)
|
||||
|
||||
def test_collect_last_token_hidden_uses_batch_owner_metadata(self):
|
||||
import torch
|
||||
@@ -506,6 +605,91 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
|
||||
self.assertEqual(collected.tolist(), [[13.0], [99.0]])
|
||||
|
||||
def test_collect_last_token_hidden_uses_compute_padding_for_single_request(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[40320],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
hidden_states = torch.zeros((64, 1), dtype=torch.float32)
|
||||
hidden_states[0] = 123.0
|
||||
hidden_states[1] = 999.0
|
||||
forward_batch = SimpleNamespace(
|
||||
extend_seq_lens_cpu=[65],
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
),
|
||||
)
|
||||
|
||||
def fake_all_gather(output, local_last):
|
||||
self.assertEqual(local_last.tolist(), [[123.0]])
|
||||
output.zero_()
|
||||
output[1] = local_last[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=1,
|
||||
),
|
||||
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, 8)
|
||||
|
||||
self.assertEqual(collected.tolist(), [[123.0]])
|
||||
|
||||
def test_collect_last_token_hidden_uses_compute_rank_offsets_for_batch(self):
|
||||
import torch
|
||||
|
||||
hidden_states = torch.zeros((8, 1), dtype=torch.float32)
|
||||
hidden_states[0] = 10.0
|
||||
hidden_states[1] = 99.0
|
||||
hidden_states[4] = 20.0
|
||||
forward_batch = SimpleNamespace(
|
||||
extend_seq_lens_cpu=[5, 5],
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=2,
|
||||
request_last_token_owner=[1, 1],
|
||||
request_last_token_local_offset=[0, 0],
|
||||
request_rank_local_offsets=[0, 1],
|
||||
request_compute_rank_local_offsets=[0, 4],
|
||||
compute_padding_enabled=True,
|
||||
),
|
||||
)
|
||||
|
||||
def fake_all_gather(output, local_last):
|
||||
self.assertEqual(local_last.tolist(), [[10.0], [20.0]])
|
||||
output.copy_(torch.tensor([[0.0], [0.0], [10.0], [20.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=1,
|
||||
),
|
||||
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(), [[10.0], [20.0]])
|
||||
|
||||
def test_collect_last_token_hidden_fails_fast_without_batch_owner_metadata(self):
|
||||
import torch
|
||||
|
||||
@@ -577,6 +761,51 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
page_size=64,
|
||||
)
|
||||
|
||||
def test_cp_shared_kv_prepare_uses_batch_plan_for_bs1_compute_padding(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=[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,
|
||||
):
|
||||
metadata = prepare_input_dp_with_cp_dsa(
|
||||
65,
|
||||
cp_rank=1,
|
||||
cp_size=8,
|
||||
seqs_len=[65],
|
||||
forward_batch=forward_batch,
|
||||
page_size=64,
|
||||
)
|
||||
|
||||
self.assertIsNotNone(metadata.batch_plan)
|
||||
self.assertEqual(metadata.batch_size, 1)
|
||||
self.assertTrue(metadata.compute_padding_enabled)
|
||||
self.assertEqual(
|
||||
metadata.request_valid_split_lists,
|
||||
[[64, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]],
|
||||
)
|
||||
self.assertEqual(
|
||||
metadata.request_compute_split_lists,
|
||||
[[64, 64, 64, 64, 64, 64, 64, 64, 0, 0, 0, 0, 0, 0, 0, 0]],
|
||||
)
|
||||
self.assertEqual(metadata.split_list, metadata.request_compute_split_lists[0])
|
||||
self.assertEqual(metadata.max_rank_len, [64] * 8)
|
||||
self.assertEqual(metadata.per_rank_actual_token, [64] * 8)
|
||||
self.assertEqual(metadata.actual_seq_q_prev, 64)
|
||||
self.assertEqual(metadata.actual_seq_q_next, 0)
|
||||
self.assertEqual(metadata.request_valid_seq_q_prev, [1])
|
||||
self.assertEqual(metadata.request_valid_seq_q_next, [0])
|
||||
|
||||
def test_cp_shared_kv_all_gather_rejects_round_robin_mode(self):
|
||||
import torch
|
||||
|
||||
@@ -792,6 +1021,34 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
|
||||
self.assertEqual(local[:, 0].tolist(), list(range(0, 16, 2)))
|
||||
|
||||
def test_cp_split_and_rebuild_data_uses_compute_padding_rows(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[40320],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
forward_batch = SimpleNamespace(
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
)
|
||||
)
|
||||
tensor = torch.arange(65 * 2, dtype=torch.float32).view(65, 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.shape, (64, 2))
|
||||
self.assertEqual(local[0].tolist(), [128.0, 129.0])
|
||||
self.assertTrue(torch.equal(local[1:], torch.zeros((63, 2))))
|
||||
|
||||
def test_cp_split_and_rebuild_1d_keeps_batch_request_boundaries(self):
|
||||
import torch
|
||||
|
||||
@@ -812,6 +1069,85 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
|
||||
self.assertEqual(local.tolist(), [8, 9, 10, 11, 12])
|
||||
|
||||
def test_cp_split_and_rebuild_1d_uses_zero_compute_padding_rows(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[40320],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
forward_batch = SimpleNamespace(
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
)
|
||||
)
|
||||
|
||||
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(65))
|
||||
|
||||
self.assertEqual(local.shape, (64,))
|
||||
self.assertEqual(local[0].item(), 64)
|
||||
self.assertEqual(local[1:].tolist(), [0] * 63)
|
||||
|
||||
def test_select_cp_local_valid_rows_filters_compute_padding_rows(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
forward_batch = SimpleNamespace(
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
)
|
||||
)
|
||||
local_compute_rows = torch.full((64, 2), -1.0)
|
||||
local_compute_rows[0] = torch.tensor([50.0, 51.0])
|
||||
|
||||
selected = select_cp_local_valid_rows_for_cache_write(
|
||||
forward_batch, local_compute_rows
|
||||
)
|
||||
|
||||
self.assertEqual(selected.tolist(), [[50.0, 51.0]])
|
||||
|
||||
def test_cp_split_and_rebuild_position_is_batch_aware_and_compute_padded(self):
|
||||
import torch
|
||||
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[40320],
|
||||
page_size=64,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
forward_batch = SimpleNamespace(
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
)
|
||||
)
|
||||
positions = torch.arange(40320, 40385, dtype=torch.int32)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
||||
return_value=False,
|
||||
):
|
||||
local = cp_split_and_rebuild_position(forward_batch, positions)
|
||||
|
||||
self.assertEqual(local.shape, (64,))
|
||||
self.assertEqual(local.tolist(), list(range(40384, 40448)))
|
||||
|
||||
def test_cp_local_embedding_pad_len_uses_metadata_max_rank_len(self):
|
||||
from types import SimpleNamespace
|
||||
|
||||
@@ -952,6 +1288,48 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
list(range(2 * page_size, 3 * page_size)) + [4 * page_size],
|
||||
)
|
||||
|
||||
def test_local_out_cache_loc_uses_valid_rows_under_compute_padding(self):
|
||||
import torch
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
out_cache_loc = torch.cat(
|
||||
[
|
||||
torch.arange(page_size, 2 * page_size),
|
||||
torch.tensor([2 * page_size]),
|
||||
]
|
||||
).to(torch.int64)
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
cp_shared_kv_layout=CpSharedKVLayout(
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
),
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
page_aligned=True,
|
||||
page_size=page_size,
|
||||
extend_prefix_len=0,
|
||||
),
|
||||
out_cache_loc=out_cache_loc,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.layers.attention.nsa.utils.is_nsa_prefill_cp_round_robin_split",
|
||||
return_value=False,
|
||||
):
|
||||
local_locs = get_cp_shared_kv_local_out_cache_loc(forward_batch)
|
||||
|
||||
self.assertEqual(local_locs.tolist(), [2 * page_size])
|
||||
|
||||
def test_batch_local_physical_out_cache_loc_reuses_layer_invariant_plan(self):
|
||||
import torch
|
||||
from types import SimpleNamespace
|
||||
@@ -1135,6 +1513,67 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
act_quant=None,
|
||||
)
|
||||
|
||||
def test_indexer_direct_write_filters_compute_padding_rows(self):
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.nsa import nsa_indexer
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
out_cache_loc = torch.cat(
|
||||
[
|
||||
torch.arange(page_size, 2 * page_size),
|
||||
torch.tensor([2 * page_size]),
|
||||
]
|
||||
).to(torch.int64)
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
cp_shared_kv_layout=CpSharedKVLayout(
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
),
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
page_aligned=True,
|
||||
page_size=page_size,
|
||||
extend_prefix_len=0,
|
||||
),
|
||||
out_cache_loc=out_cache_loc,
|
||||
token_to_kv_pool=SimpleNamespace(page_size=page_size),
|
||||
)
|
||||
indexer = object.__new__(Indexer)
|
||||
indexer.nsa_enable_prefill_cp = True
|
||||
calls = []
|
||||
|
||||
def fake_store_index_k_cache(**kwargs):
|
||||
calls.append(kwargs)
|
||||
|
||||
indexer._store_index_k_cache = fake_store_index_k_cache
|
||||
local_key = torch.full((64, 2), -1.0)
|
||||
local_key[0] = torch.tensor([9.0, 10.0])
|
||||
|
||||
with patch.object(nsa_indexer, "nsa_use_prefill_cp", return_value=True):
|
||||
stored = Indexer._store_cp_shared_local_index_k_cache(
|
||||
indexer,
|
||||
forward_batch,
|
||||
layer_id=0,
|
||||
local_key=local_key,
|
||||
act_quant=None,
|
||||
)
|
||||
|
||||
self.assertTrue(stored)
|
||||
self.assertEqual(len(calls), 1)
|
||||
self.assertEqual(calls[0]["key"].tolist(), [[9.0, 10.0]])
|
||||
|
||||
def test_mla_direct_write_fails_fast_on_local_shape_mismatch(self):
|
||||
import torch
|
||||
|
||||
@@ -1161,6 +1600,292 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
k_pe=torch.empty((2, 8)),
|
||||
)
|
||||
|
||||
def test_mla_direct_write_filters_compute_padding_rows(self):
|
||||
import torch
|
||||
|
||||
from sglang.srt.models.deepseek_common.attention_forward_methods import (
|
||||
forward_mla,
|
||||
)
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
out_cache_loc = torch.cat(
|
||||
[
|
||||
torch.arange(page_size, 2 * page_size),
|
||||
torch.tensor([2 * page_size]),
|
||||
]
|
||||
).to(torch.int64)
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
cp_shared_kv_layout=CpSharedKVLayout(
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
),
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
page_aligned=True,
|
||||
page_size=page_size,
|
||||
extend_prefix_len=0,
|
||||
),
|
||||
out_cache_loc=out_cache_loc,
|
||||
token_to_kv_pool=SimpleNamespace(page_size=page_size),
|
||||
)
|
||||
mla = SimpleNamespace(attn_mqa=SimpleNamespace(layer_id=0))
|
||||
calls = []
|
||||
|
||||
def fake_tai_store(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return True
|
||||
|
||||
k_nope = torch.full((64, 2), -1.0)
|
||||
k_nope[0] = torch.tensor([1.0, 2.0])
|
||||
k_pe = torch.full((64, 2), -1.0)
|
||||
k_pe[0] = torch.tensor([3.0, 4.0])
|
||||
|
||||
with patch.object(forward_mla, "try_tai_fused_mla_store", fake_tai_store):
|
||||
stored = (
|
||||
forward_mla.DeepseekMLAForwardMixin._maybe_write_cp_shared_local_mla_kv(
|
||||
mla,
|
||||
forward_batch,
|
||||
k_nope=k_nope,
|
||||
k_pe=k_pe,
|
||||
)
|
||||
)
|
||||
|
||||
self.assertTrue(stored)
|
||||
self.assertEqual(len(calls), 1)
|
||||
self.assertEqual(calls[0]["k_nope"].tolist(), [[1.0, 2.0]])
|
||||
self.assertEqual(calls[0]["k_rope"].tolist(), [[3.0, 4.0]])
|
||||
self.assertEqual(calls[0]["logical_locs"].tolist(), [2 * page_size])
|
||||
|
||||
def test_index_partial_current_compose_accepts_local_valid_compute_padding_rows(
|
||||
self,
|
||||
):
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.nsa import nsa_indexer
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
|
||||
class FakePool:
|
||||
page_size = 64
|
||||
index_head_dim = 2
|
||||
|
||||
def get_index_k_with_scale_buffer(self, layer_id):
|
||||
return torch.zeros((4, 3), dtype=torch.float32)
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
cp_shared_kv_layout=CpSharedKVLayout(
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
),
|
||||
cp_shared_kv_index_prefetcher=None,
|
||||
token_to_kv_pool=FakePool(),
|
||||
nsa_cp_metadata=NSAContextParallelMetadata(
|
||||
batch_size=1,
|
||||
batch_plan=plan,
|
||||
page_aligned=True,
|
||||
page_size=page_size,
|
||||
extend_prefix_len=0,
|
||||
),
|
||||
out_cache_loc=torch.cat(
|
||||
[
|
||||
torch.arange(page_size, 2 * page_size),
|
||||
torch.tensor([2 * page_size]),
|
||||
]
|
||||
).to(torch.int64),
|
||||
extend_prefix_lens_cpu=[page_size],
|
||||
extend_seq_lens_cpu=[65],
|
||||
)
|
||||
logical_page_table = torch.tensor([[1, 2, 3]], dtype=torch.int32)
|
||||
current_index_kv = (
|
||||
torch.tensor([[7.0, 8.0]], dtype=torch.float32),
|
||||
torch.tensor([[0.5]], dtype=torch.float32),
|
||||
)
|
||||
materialize_calls = []
|
||||
expected_buffer = torch.ones((3, 3), dtype=torch.float32)
|
||||
expected_pages = torch.tensor([[0, 1, 2]], dtype=torch.int32)
|
||||
indexer = object.__new__(Indexer)
|
||||
|
||||
def fake_materialize(**kwargs):
|
||||
materialize_calls.append(kwargs)
|
||||
return expected_buffer, expected_pages
|
||||
|
||||
with patch.object(
|
||||
nsa_indexer,
|
||||
"get_or_build_shared_paged_buffer_slot_remap",
|
||||
return_value=torch.tensor([0, 1, 2], dtype=torch.int64),
|
||||
), patch.object(
|
||||
nsa_indexer,
|
||||
"materialize_prefix_and_reuse_current_index_page_slots",
|
||||
side_effect=fake_materialize,
|
||||
):
|
||||
dense_buffer, dense_pages = indexer._maybe_materialize_shared_index_buffer(
|
||||
forward_batch,
|
||||
layer_id=0,
|
||||
logical_page_table=logical_page_table,
|
||||
current_index_kv=current_index_kv,
|
||||
)
|
||||
|
||||
self.assertIs(dense_buffer, expected_buffer)
|
||||
self.assertIs(dense_pages, expected_pages)
|
||||
self.assertEqual(len(materialize_calls), 1)
|
||||
self.assertIs(materialize_calls[0]["current_index_k"], current_index_kv[0])
|
||||
self.assertEqual(materialize_calls[0]["current_locs"].tolist(), [2 * page_size])
|
||||
|
||||
def test_indexer_current_reuse_compute_padding_selects_local_key_not_gathered_key(
|
||||
self,
|
||||
):
|
||||
import torch
|
||||
import types
|
||||
|
||||
from sglang.srt.layers.attention.nsa import nsa_indexer
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
metadata_obj = _build_batch_metadata_from_plan(plan)
|
||||
|
||||
class Mode:
|
||||
def is_extend_without_speculative(self):
|
||||
return True
|
||||
|
||||
def is_decode_or_idle(self):
|
||||
return False
|
||||
|
||||
def is_target_verify(self):
|
||||
return False
|
||||
|
||||
def is_draft_extend(self, include_v2=False):
|
||||
return False
|
||||
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
|
||||
class AttnBackend:
|
||||
def get_indexer_metadata(self, layer_id, forward_batch):
|
||||
return object()
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
cp_shared_kv_layout=CpSharedKVLayout(
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
),
|
||||
token_to_kv_pool=SimpleNamespace(page_size=page_size),
|
||||
nsa_cp_metadata=metadata_obj,
|
||||
out_cache_loc=torch.cat(
|
||||
[
|
||||
torch.arange(page_size, 2 * page_size),
|
||||
torch.tensor([2 * page_size]),
|
||||
]
|
||||
).to(torch.int64),
|
||||
extend_prefix_lens_cpu=[0],
|
||||
extend_seq_lens_cpu=[65],
|
||||
seq_lens_cpu=torch.tensor([65], dtype=torch.int64),
|
||||
forward_mode=Mode(),
|
||||
attn_backend=AttnBackend(),
|
||||
hisparse_coordinator=None,
|
||||
)
|
||||
|
||||
local_key = torch.full((64, 2), -1.0, dtype=torch.float32)
|
||||
local_key[0] = torch.tensor([11.0, 12.0])
|
||||
gathered_key = torch.full((64, 2), 99.0, dtype=torch.float32)
|
||||
gathered_key[0] = torch.tensor([101.0, 102.0])
|
||||
query = torch.zeros((64, 2), dtype=torch.float32)
|
||||
act_quant_inputs = []
|
||||
topk_current_index_kv = []
|
||||
|
||||
def fake_act_quant(tensor, block_size, scale_fmt):
|
||||
act_quant_inputs.append(tensor.detach().clone())
|
||||
return tensor.detach().clone(), torch.ones(
|
||||
(int(tensor.shape[0]), 1), dtype=torch.float32
|
||||
)
|
||||
|
||||
fake_triton_kernel = types.ModuleType(
|
||||
"sglang.srt.layers.attention.nsa.triton_kernel"
|
||||
)
|
||||
fake_triton_kernel.act_quant = fake_act_quant
|
||||
|
||||
indexer = object.__new__(Indexer)
|
||||
indexer.alt_stream = None
|
||||
indexer.nsa_enable_prefill_cp = True
|
||||
indexer.index_topk = 2
|
||||
indexer.block_size = 64
|
||||
indexer.scale_fmt = None
|
||||
indexer._get_q_k_bf16 = (
|
||||
lambda *args, **kwargs: (query, gathered_key, local_key)
|
||||
)
|
||||
indexer._store_cp_shared_local_index_k_cache = lambda **kwargs: True
|
||||
indexer._can_reuse_current_index_kv = lambda forward_batch: True
|
||||
indexer._get_logits_head_gate = (
|
||||
lambda x_for_gate, q_scale: torch.zeros((64, 1), dtype=torch.float32)
|
||||
)
|
||||
|
||||
def fake_topk(*args, **kwargs):
|
||||
topk_current_index_kv.append(kwargs["current_index_kv"])
|
||||
return torch.zeros((64, 2), dtype=torch.int32)
|
||||
|
||||
indexer._get_topk_in_seq_cp_pair = fake_topk
|
||||
|
||||
with (
|
||||
patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"sglang.srt.layers.attention.nsa.triton_kernel": fake_triton_kernel
|
||||
},
|
||||
),
|
||||
patch.object(nsa_indexer, "_is_cuda", True),
|
||||
patch.object(nsa_indexer, "_is_hip", False),
|
||||
patch.object(nsa_indexer, "_is_npu", False),
|
||||
patch.object(
|
||||
nsa_indexer,
|
||||
"is_nsa_prefill_cp_in_seq_split",
|
||||
return_value=True,
|
||||
),
|
||||
):
|
||||
result = Indexer.forward_cuda(
|
||||
indexer,
|
||||
x=torch.zeros((64, 2), dtype=torch.float32),
|
||||
q_lora=torch.zeros((64, 2), dtype=torch.float32),
|
||||
positions=torch.arange(64, dtype=torch.int64),
|
||||
forward_batch=forward_batch,
|
||||
layer_id=0,
|
||||
return_indices=True,
|
||||
)
|
||||
|
||||
self.assertEqual(result.shape, (64, 2))
|
||||
self.assertEqual(len(act_quant_inputs), 2)
|
||||
self.assertEqual(act_quant_inputs[1].tolist(), [[11.0, 12.0]])
|
||||
self.assertNotEqual(act_quant_inputs[1].tolist(), [[101.0, 102.0]])
|
||||
self.assertEqual(len(topk_current_index_kv), 1)
|
||||
self.assertEqual(topk_current_index_kv[0][0].tolist(), [[11.0, 12.0]])
|
||||
|
||||
def test_indexer_direct_write_does_not_log_missing_metadata_for_non_cp_batch(self):
|
||||
import torch
|
||||
from types import SimpleNamespace
|
||||
@@ -1416,6 +2141,113 @@ class TestNSAInSeqCPUtils(unittest.TestCase):
|
||||
[[1, 1], [2, 2], [3, 3], [4, 4], [5, 5], [6, 6], [7, 7]],
|
||||
)
|
||||
|
||||
def test_indexer_in_seq_cp_pair_compute_padding_outputs_dummy_safe_rows(self):
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.nsa.nsa_indexer import Indexer
|
||||
|
||||
page_size = 64
|
||||
plan = build_batch_page_aligned_in_seq_split_plan(
|
||||
extend_lens=[65],
|
||||
prefix_lens=[0],
|
||||
page_size=page_size,
|
||||
cp_size=8,
|
||||
cp_rank=1,
|
||||
)
|
||||
metadata_obj = _build_batch_metadata_from_plan(plan)
|
||||
indexer = object.__new__(Indexer)
|
||||
indexer.index_topk = 2
|
||||
logical_pages = torch.tensor([[1, 2]], dtype=torch.int32)
|
||||
materialized_index = torch.tensor([11], dtype=torch.int32)
|
||||
dense_pages = torch.tensor([[1, 2]], dtype=torch.int32)
|
||||
materialize_calls = []
|
||||
topk_calls = []
|
||||
|
||||
class Metadata:
|
||||
def get_page_table_64(self):
|
||||
return logical_pages
|
||||
|
||||
def get_page_table_1(self):
|
||||
return torch.empty((1, 65), dtype=torch.int32)
|
||||
|
||||
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,
|
||||
batch_idx=0,
|
||||
):
|
||||
topk_calls.append(
|
||||
{
|
||||
"actual_seq_q": actual_seq_q,
|
||||
"cp_index": cp_index,
|
||||
"q": q_fp8.flatten().tolist(),
|
||||
"weights": weights.flatten().tolist(),
|
||||
"shared_index_buffer": shared_index_buffer,
|
||||
"shared_block_tables": shared_block_tables,
|
||||
}
|
||||
)
|
||||
rows = int(q_fp8.shape[0])
|
||||
return (
|
||||
torch.arange(1, rows + 1, dtype=torch.int32)
|
||||
.view(rows, 1)
|
||||
.repeat(1, 2)
|
||||
)
|
||||
|
||||
indexer._maybe_materialize_shared_index_buffer = fake_materialize
|
||||
indexer._get_topk_ragged_with_cp = fake_get_topk
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
batch_size=1,
|
||||
forward_mode=SimpleNamespace(
|
||||
is_extend_without_speculative=lambda: True,
|
||||
),
|
||||
extend_prefix_lens_cpu=[0],
|
||||
extend_seq_lens_cpu=[65],
|
||||
seq_lens_cpu=torch.tensor([65], dtype=torch.int64),
|
||||
nsa_cp_metadata=metadata_obj,
|
||||
)
|
||||
q_fp8 = torch.arange(64, dtype=torch.float32).view(64, 1)
|
||||
weights = (torch.arange(64, dtype=torch.float32) + 100).view(64, 1)
|
||||
|
||||
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), 1)
|
||||
self.assertEqual(topk_calls[0]["actual_seq_q"], 1)
|
||||
self.assertEqual(topk_calls[0]["cp_index"], [(0, 64, 65)])
|
||||
self.assertEqual(topk_calls[0]["q"], [0.0])
|
||||
self.assertEqual(topk_calls[0]["weights"], [100.0])
|
||||
self.assertIs(topk_calls[0]["shared_index_buffer"], materialized_index)
|
||||
self.assertIs(topk_calls[0]["shared_block_tables"], dense_pages)
|
||||
self.assertEqual(result.shape, (64, 2))
|
||||
self.assertEqual(result[0].tolist(), [1, 1])
|
||||
self.assertTrue(
|
||||
torch.equal(result[1:], torch.full((63, 2), -1, dtype=torch.int32))
|
||||
)
|
||||
|
||||
def test_indexer_in_seq_cp_pair_batch_materializes_partial_current_index_reuse_once(self):
|
||||
import torch
|
||||
|
||||
|
||||
@@ -297,6 +297,173 @@ class TestCPSharedPagedAllocator(CustomTestCase):
|
||||
|
||||
self.assertEqual(owners, [0, 1])
|
||||
|
||||
def test_alloc_extend_compute_owner_uses_valid_pages_not_compute_padding_pages(
|
||||
self,
|
||||
):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.mem_cache import common
|
||||
|
||||
page_size = 64
|
||||
|
||||
class FakeAllocator:
|
||||
def __init__(self):
|
||||
self.page_size = page_size
|
||||
self.cp_size = 8
|
||||
self.owner_calls = []
|
||||
self.extend_num_tokens = []
|
||||
|
||||
def alloc_extend_compute_owner(
|
||||
self,
|
||||
_prefix_lens,
|
||||
_prefix_lens_cpu,
|
||||
_seq_lens,
|
||||
_seq_lens_cpu,
|
||||
_last_loc,
|
||||
extend_num_tokens,
|
||||
page_compute_owners,
|
||||
):
|
||||
self.extend_num_tokens.append(int(extend_num_tokens))
|
||||
self.owner_calls.append(list(page_compute_owners))
|
||||
return torch.arange(
|
||||
1024, 1024 + int(extend_num_tokens), dtype=torch.int64
|
||||
)
|
||||
|
||||
def alloc_extend(self, *_args, **_kwargs):
|
||||
raise AssertionError("legacy allocation should not be used")
|
||||
|
||||
class FakeTreeCache:
|
||||
def __init__(self):
|
||||
self.token_to_kv_pool_allocator = FakeAllocator()
|
||||
|
||||
def is_chunk_cache(self):
|
||||
return False
|
||||
|
||||
def evict(self, *_args, **_kwargs):
|
||||
raise AssertionError("eviction should not be needed")
|
||||
|
||||
server_args = SimpleNamespace(
|
||||
enable_nsa_prefill_cp_shared_kv=True,
|
||||
enable_nsa_prefill_context_parallel=True,
|
||||
nsa_prefill_cp_mode="in-seq-split",
|
||||
)
|
||||
tree_cache = FakeTreeCache()
|
||||
|
||||
with patch.object(common, "get_global_server_args", return_value=server_args):
|
||||
out_cache_loc = common.alloc_paged_token_slots_extend(
|
||||
tree_cache=tree_cache,
|
||||
prefix_lens=torch.tensor([0], dtype=torch.int64),
|
||||
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
|
||||
seq_lens=torch.tensor([65], dtype=torch.int64),
|
||||
seq_lens_cpu=torch.tensor([65], dtype=torch.int64),
|
||||
last_loc=torch.tensor([-1], dtype=torch.int64),
|
||||
extend_num_tokens=65,
|
||||
)
|
||||
|
||||
self.assertEqual(out_cache_loc.numel(), 65)
|
||||
self.assertEqual(tree_cache.token_to_kv_pool_allocator.extend_num_tokens, [65])
|
||||
self.assertEqual(tree_cache.token_to_kv_pool_allocator.owner_calls, [[0, 1]])
|
||||
|
||||
def test_cp_hicache_write_reservation_uses_page_tail_not_compute_padding_extent(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.managers.cache_controller import HiCacheController
|
||||
|
||||
page_size = 64
|
||||
|
||||
class FakeHostPool:
|
||||
def __init__(self):
|
||||
self.alloc_sizes = []
|
||||
|
||||
def alloc_contiguous_preferred(self, need_size):
|
||||
self.alloc_sizes.append(int(need_size))
|
||||
return torch.arange(1024, 1024 + int(need_size), dtype=torch.int64)
|
||||
|
||||
def alloc(self, need_size):
|
||||
return self.alloc_contiguous_preferred(need_size)
|
||||
|
||||
def free(self, _indices):
|
||||
raise AssertionError("reservation should not roll back")
|
||||
|
||||
controller = HiCacheController.__new__(HiCacheController)
|
||||
controller.page_size = page_size
|
||||
controller.cp_shared_kv_layout = CpSharedKVLayout(
|
||||
page_size=page_size, cp_size=8, cp_rank=1
|
||||
)
|
||||
controller.mem_pool_host = FakeHostPool()
|
||||
controller.draft_mem_pool_host = None
|
||||
controller.draft_mem_pool_device = None
|
||||
|
||||
reservation = controller.reserve_write_cp(
|
||||
torch.arange(page_size, page_size + 65, dtype=torch.int64),
|
||||
node_id=123,
|
||||
)
|
||||
|
||||
self.assertEqual(controller.mem_pool_host.alloc_sizes, [page_size])
|
||||
self.assertEqual(reservation.metadata.logical_len, 65)
|
||||
self.assertEqual(reservation.metadata.padded_len, page_size * 2)
|
||||
self.assertEqual(reservation.metadata.page_owners.tolist(), [0, 1])
|
||||
self.assertEqual(reservation.host_indices.numel(), page_size)
|
||||
self.assertEqual(reservation.physical_device_indices.numel(), page_size)
|
||||
self.assertEqual(
|
||||
reservation.metadata.owned_positions.tolist(),
|
||||
list(range(64, 128)),
|
||||
)
|
||||
|
||||
def test_cp_hicache_load_returns_valid_visible_len_while_loading_owned_page_tail(
|
||||
self,
|
||||
):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sglang.srt.managers.cache_controller import HiCacheController
|
||||
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
|
||||
|
||||
page_size = 64
|
||||
|
||||
class FakeDeviceAllocator:
|
||||
def __init__(self):
|
||||
self.owner_calls = []
|
||||
self.freed = []
|
||||
|
||||
def alloc_pages_with_owners(self, page_owners):
|
||||
self.owner_calls.append(list(page_owners))
|
||||
return torch.arange(page_size, page_size * 3, dtype=torch.int64)
|
||||
|
||||
def free(self, indices):
|
||||
self.freed.append(indices.clone())
|
||||
|
||||
controller = HiCacheController.__new__(HiCacheController)
|
||||
controller.page_size = page_size
|
||||
controller.cp_shared_kv_layout = CpSharedKVLayout(
|
||||
page_size=page_size, cp_size=8, cp_rank=1
|
||||
)
|
||||
controller.mem_pool_device_allocator = FakeDeviceAllocator()
|
||||
controller.load_queue = []
|
||||
controller.draft_load_queue = []
|
||||
controller.draft_mem_pool_host = None
|
||||
controller.draft_mem_pool_device = None
|
||||
|
||||
metadata = CpHiCacheNodeMetadata(
|
||||
logical_len=65,
|
||||
padded_len=page_size * 2,
|
||||
owned_positions=torch.arange(page_size, page_size * 2, dtype=torch.int64),
|
||||
host_indices=torch.arange(1024, 1024 + page_size, dtype=torch.int64),
|
||||
page_owners=torch.tensor([0, 1], dtype=torch.int8),
|
||||
page_size=page_size,
|
||||
)
|
||||
node = SimpleNamespace(cp_hicache=metadata, host_len=65, id=321)
|
||||
|
||||
visible_device_indices = controller.load_cp([node], node_id=321)
|
||||
|
||||
self.assertEqual(controller.mem_pool_device_allocator.owner_calls, [[0, 1]])
|
||||
self.assertEqual(controller.mem_pool_device_allocator.freed, [])
|
||||
self.assertEqual(visible_device_indices.numel(), 65)
|
||||
self.assertEqual(visible_device_indices.tolist(), list(range(64, 129)))
|
||||
self.assertEqual(len(controller.load_queue), 1)
|
||||
load_op = controller.load_queue[0]
|
||||
self.assertEqual(load_op.host_indices.tolist(), list(range(1024, 1024 + 64)))
|
||||
self.assertEqual(load_op.device_indices.tolist(), list(range(64, 128)))
|
||||
|
||||
def test_compute_owner_page_assignment_allows_radix_hit_suffix_with_one_page_per_rank(
|
||||
self,
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user