[Feature] Support JIT set kv cache (#16273)
This commit is contained in:
35
python/sglang/jit_kernel/tests/test_store_cache.py
Normal file
35
python/sglang/jit_kernel/tests/test_store_cache.py
Normal file
@@ -0,0 +1,35 @@
|
||||
import itertools
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.kvcache import store_cache
|
||||
|
||||
BS_LIST = [2**n for n in range(0, 15)]
|
||||
BS_LIST += [x + 1 + i for i, x in enumerate(BS_LIST)]
|
||||
HIDDEN_DIMS = [64, 128, 256, 512, 1024, 96, 98, 100]
|
||||
CACHE_SIZE = 1024 * 1024
|
||||
DTYPE = torch.bfloat16
|
||||
DEVICE = "cuda"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"batch_size,element_dim",
|
||||
list(itertools.product(BS_LIST, HIDDEN_DIMS)),
|
||||
)
|
||||
def test_store_cache(batch_size: int, element_dim: int) -> None:
|
||||
k = torch.randn((batch_size, element_dim), dtype=DTYPE, device=DEVICE)
|
||||
v = torch.randn((batch_size, element_dim), dtype=DTYPE, device=DEVICE)
|
||||
k_cache = torch.randn((CACHE_SIZE, element_dim), dtype=DTYPE, device=DEVICE)
|
||||
v_cache = torch.randn((CACHE_SIZE, element_dim), dtype=DTYPE, device=DEVICE)
|
||||
indices = torch.randperm(CACHE_SIZE, device=DEVICE)[:batch_size]
|
||||
|
||||
# AOT store cache
|
||||
store_cache(k, v, k_cache, v_cache, indices)
|
||||
|
||||
assert torch.all(k_cache[indices] == k)
|
||||
assert torch.all(v_cache[indices] == v)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
Reference in New Issue
Block a user