Files
sglang/test/registered/unit/mem_cache/test_cp_hicache_metadata.py
T

144 lines
5.4 KiB
Python

import sys
import unittest
from unittest.mock import MagicMock
import torch
# Stub out sgl_kernel before any sglang import so this CPU unit test does not
# require CUDA extension libraries to be installed.
for _mod in ("sgl_kernel", "sgl_kernel.kvcacheio"):
if _mod not in sys.modules:
sys.modules[_mod] = MagicMock()
from sglang.srt.mem_cache.hiradix_cache import CpHiCacheNodeMetadata
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=2, suite="stage-a-test-cpu")
class TestCpHiCacheNodeMetadata(CustomTestCase):
def test_split_zero_len_moves_all_positions_to_child(self):
metadata = CpHiCacheNodeMetadata(
logical_len=8,
owned_positions=torch.tensor([1, 3, 7], dtype=torch.int64),
host_indices=torch.tensor([10, 11, 12], dtype=torch.int64),
)
parent, child = metadata.split(0)
self.assertEqual(parent.logical_len, 0)
self.assertEqual(parent.owned_positions.tolist(), [])
self.assertEqual(parent.host_indices.tolist(), [])
self.assertEqual(child.logical_len, 8)
self.assertEqual(child.owned_positions.tolist(), [1, 3, 7])
self.assertEqual(child.host_indices.tolist(), [10, 11, 12])
def test_split_rebases_child_positions(self):
metadata = CpHiCacheNodeMetadata(
logical_len=10,
owned_positions=torch.tensor([0, 2, 5, 9], dtype=torch.int64),
host_indices=torch.tensor([20, 21, 22, 23], dtype=torch.int64),
)
parent, child = metadata.split(5)
self.assertEqual(parent.logical_len, 5)
self.assertEqual(parent.owned_positions.tolist(), [0, 2])
self.assertEqual(parent.host_indices.tolist(), [20, 21])
self.assertEqual(child.logical_len, 5)
self.assertEqual(child.owned_positions.tolist(), [0, 4])
self.assertEqual(child.host_indices.tolist(), [22, 23])
def test_zero_owned_metadata_is_valid(self):
metadata = CpHiCacheNodeMetadata(
logical_len=64,
owned_positions=torch.empty((0,), dtype=torch.int32),
host_indices=torch.empty((0,), dtype=torch.int32),
)
self.assertEqual(metadata.logical_len, 64)
self.assertEqual(metadata.owned_positions.device.type, "cpu")
self.assertEqual(metadata.host_indices.device.type, "cpu")
self.assertEqual(metadata.owned_positions.dtype, torch.int64)
self.assertEqual(metadata.host_indices.dtype, torch.int64)
def test_non_int64_inputs_are_converted(self):
metadata = CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([1, 3], dtype=torch.int32),
host_indices=torch.tensor([10, 11], dtype=torch.int32),
)
self.assertEqual(metadata.owned_positions.dtype, torch.int64)
self.assertEqual(metadata.host_indices.dtype, torch.int64)
def test_negative_logical_len_raises(self):
with self.assertRaisesRegex(ValueError, "logical_len"):
CpHiCacheNodeMetadata(
logical_len=-1,
owned_positions=torch.empty((0,), dtype=torch.int64),
host_indices=torch.empty((0,), dtype=torch.int64),
)
def test_metadata_does_not_alias_input_tensors(self):
owned_positions = torch.tensor([1, 3], dtype=torch.int64)
host_indices = torch.tensor([10, 11], dtype=torch.int64)
metadata = CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=owned_positions,
host_indices=host_indices,
)
owned_positions[0] = 2
host_indices[0] = 12
self.assertEqual(metadata.owned_positions.tolist(), [1, 3])
self.assertEqual(metadata.host_indices.tolist(), [10, 11])
def test_invalid_split_raises(self):
metadata = CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([1], dtype=torch.int64),
host_indices=torch.tensor([9], dtype=torch.int64),
)
with self.assertRaisesRegex(ValueError, "split_len"):
metadata.split(5)
def test_unsorted_positions_raise(self):
with self.assertRaisesRegex(ValueError, "sorted"):
CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([2, 1], dtype=torch.int64),
host_indices=torch.tensor([9, 10], dtype=torch.int64),
)
def test_duplicate_positions_raise(self):
with self.assertRaisesRegex(ValueError, "strictly increasing"):
CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([1, 1], dtype=torch.int64),
host_indices=torch.tensor([9, 10], dtype=torch.int64),
)
def test_length_mismatch_raises(self):
with self.assertRaisesRegex(ValueError, "same length"):
CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([1, 2], dtype=torch.int64),
host_indices=torch.tensor([9], dtype=torch.int64),
)
def test_out_of_range_positions_raise(self):
with self.assertRaisesRegex(ValueError, r"\[0, logical_len\)"):
CpHiCacheNodeMetadata(
logical_len=4,
owned_positions=torch.tensor([4], dtype=torch.int64),
host_indices=torch.tensor([9], dtype=torch.int64),
)
if __name__ == "__main__":
unittest.main()