[XPU]Enable HiSparse hierarchical sparse KV cache on Intel XPU (#32792)

Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
AMRUTHA M
2026-09-21 14:07:49 +08:00
committed by GitHub
co-authored by Ma Mingfei
parent fcb080bd40
commit d20cd9d77f
7 changed files with 126 additions and 65 deletions
@@ -3,26 +3,43 @@ import sys
import pytest
import torch
from sglang.kernels.ops.kvcache.hisparse import (
load_cache_to_device_buffer_dsv4_mla,
load_cache_to_device_buffer_mla,
transfer_cache_dsv4_mla,
from sglang.srt.utils import (
get_device,
get_device_module,
is_cuda,
is_hip,
is_xpu,
)
from sglang.test.ci.ci_register import (
register_amd_ci,
register_cuda_ci,
register_xpu_ci,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
register_amd_ci(est_time=30, stage="stage-b", runner_config="1-gpu-small-amd")
register_cuda_ci(est_time=10, stage="base-b-kernel-unit", runner_config="1-gpu-large")
register_xpu_ci(est_time=60, suite="stage-b-test-1-gpu-xpu")
if is_xpu():
from sgl_kernel import (
load_cache_to_device_buffer_dsv4_mla,
load_cache_to_device_buffer_mla,
transfer_cache_dsv4_mla,
)
else:
from sglang.kernels.ops.kvcache.hisparse import (
load_cache_to_device_buffer_dsv4_mla,
load_cache_to_device_buffer_mla,
transfer_cache_dsv4_mla,
)
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available()
or is_npu()
or is_xpu()
or not (is_cuda() or is_hip()),
reason="HiSparse JIT tests require CUDA/ROCm.",
not (is_cuda() or is_hip() or is_xpu()),
reason="HiSparse kernel tests require CUDA/ROCm/XPU.",
)
DEVICE = "cuda"
DEVICE = get_device()
DTYPE = torch.float32
KV_DIM = 8
HOT_BUFFER_SIZE = 4
@@ -131,7 +148,7 @@ def _run_kernel(
block_size=256,
num_real_reqs=torch.tensor([num_real_reqs], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
return out
@@ -171,7 +188,7 @@ def _make_state(
device_buffer[device_buffer_locs[rid, HOT_BUFFER_SIZE]].copy_(
host_cache[newest_token].to(DEVICE, non_blocking=True)
)
torch.cuda.synchronize()
get_device_module().synchronize()
return {
"host_cache": host_cache,
@@ -199,7 +216,7 @@ def test_transfer_cache_dsv4_mla_copies_paged_token() -> None:
src_indices=torch.tensor([src_loc], dtype=torch.int64, device=DEVICE),
dst_indices=torch.tensor([dst_loc], dtype=torch.int64, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
assert torch.equal(
_read_dsv4_token(dst_cache, dst_loc).to(DEVICE),
@@ -251,7 +268,7 @@ def test_dsv4_swap_in_reads_paged_host_layout() -> None:
block_size=256,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
assert out.item() == swap_loc
assert torch.equal(
@@ -427,7 +444,7 @@ def test_load_cache_to_device_buffer_miss_copy_is_byte_exact(
)
for slot in range(HOT_BUFFER_SIZE):
device_buffer[slot].copy_(host_cache[slot].to(DEVICE))
torch.cuda.synchronize()
get_device_module().synchronize()
top_k_tokens = torch.tensor([[miss_token]], dtype=torch.int32, device=DEVICE)
out = torch.full_like(top_k_tokens, -1)
@@ -454,7 +471,7 @@ def test_load_cache_to_device_buffer_miss_copy_is_byte_exact(
block_size=256,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
# The miss evicts the LRU head (slot 0, physical loc 0) and lands there.
assert torch.equal(out.cpu(), torch.tensor([[0]], dtype=torch.int32))
@@ -593,7 +610,7 @@ def test_load_cache_to_device_buffer_dsv4_mla_miss_copy_layout() -> None:
block_size=256,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
assert torch.equal(out.cpu(), torch.tensor([[9]], dtype=torch.int32))
@@ -666,7 +683,7 @@ def test_load_cache_to_device_buffer_dsv4_fused_copy_multi_miss() -> None:
block_size=256,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
# Which slot each miss evicts is up to the LRU, so take the destinations
# from the kernel; only require that they are distinct and in range.
@@ -741,7 +758,7 @@ def test_load_cache_to_device_buffer_rocm_large_lru_writeback() -> None:
block_size=1024,
num_real_reqs=torch.tensor([1], dtype=torch.int32, device=DEVICE),
)
torch.cuda.synchronize()
get_device_module().synchronize()
expected_lru = torch.cat(
[
@@ -17,12 +17,26 @@ import torch
from sglang.srt.managers.schedule_batch import ReqKvInfo
from sglang.srt.runtime_context import publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.srt.utils import (
get_device,
get_device_module,
is_cuda,
is_hip,
is_xpu,
)
from sglang.srt.utils.common import Range
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.ci.ci_register import (
register_amd_ci,
register_cuda_ci,
register_xpu_ci,
)
register_cuda_ci(est_time=11, stage="base-b", runner_config="1-gpu-small")
register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd")
register_xpu_ci(est_time=60, suite="stage-b-test-1-gpu-xpu")
DEVICE = get_device()
# ---------------------------------------------------------------------------
# Test configuration (small-scale for fast CI runs)
@@ -74,12 +88,8 @@ class TestHiSparseUnit(unittest.TestCase):
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA is required for HiSparse tests.")
if is_npu() or is_xpu():
raise unittest.SkipTest("HiSparse tests only support CUDA/ROCm.")
if not (is_cuda() or is_hip()):
raise unittest.SkipTest("CUDA/ROCm not available.")
if not (is_cuda() or is_hip() or is_xpu()):
raise unittest.SkipTest("CUDA/ROCm/XPU not available.")
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29599")
@@ -92,8 +102,8 @@ class TestHiSparseUnit(unittest.TestCase):
alloc_with_pin_memory,
)
cls._original_alloc = ALLOC_MEMORY_FUNCS["cuda"]
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
cls._original_alloc = ALLOC_MEMORY_FUNCS[DEVICE]
ALLOC_MEMORY_FUNCS[DEVICE] = alloc_with_pin_memory
if is_hip():
from sglang.srt.layers.attention.dsa.utils import (
@@ -116,7 +126,7 @@ class TestHiSparseUnit(unittest.TestCase):
dtype=torch.bfloat16,
qk_rope_head_dim=QK_ROPE_HEAD_DIM,
layer_num=LAYER_NUM,
device="cuda",
device=DEVICE,
index_head_dim=128,
enable_memory_saver=False,
kv_cache_dim=KV_CACHE_DIM,
@@ -126,7 +136,7 @@ class TestHiSparseUnit(unittest.TestCase):
size=SIZE,
page_size=global_page_size,
dtype=torch.bfloat16,
device="cuda",
device=DEVICE,
kvcache=cls.device_pool,
need_sort=False,
host_to_device_ratio=HOST_TO_DEVICE_RATIO,
@@ -137,7 +147,7 @@ class TestHiSparseUnit(unittest.TestCase):
cls.req_to_token_pool = ReqToTokenPool(
size=MAX_NUM_REQS,
max_context_len=MAX_CONTEXT_LEN,
device="cuda",
device=DEVICE,
enable_memory_saver=False,
)
@@ -149,7 +159,7 @@ class TestHiSparseUnit(unittest.TestCase):
token_to_kv_pool_allocator=cls.allocator,
top_k=TOP_K,
device_buffer_size=DEVICE_BUFFER_SIZE,
device="cuda",
device=DEVICE,
tp_group=cls.tp_group,
host_to_device_ratio=HOST_TO_DEVICE_RATIO,
)
@@ -158,7 +168,7 @@ class TestHiSparseUnit(unittest.TestCase):
def tearDownClass(cls):
from sglang.srt.mem_cache.pool_host.common import ALLOC_MEMORY_FUNCS
ALLOC_MEMORY_FUNCS["cuda"] = cls._original_alloc
ALLOC_MEMORY_FUNCS[DEVICE] = cls._original_alloc
if torch.distributed.is_initialized():
torch.distributed.destroy_process_group()
@@ -260,11 +270,11 @@ class TestHiSparseUnit(unittest.TestCase):
def _populate_host_pool(self, req, fill_len):
"""Allocate host slots, write known patterns, register in coordinator.
Returns host_indices (cuda tensor)."""
Returns host_indices (device tensor)."""
host_pool = self.coordinator.mem_pool_host
host_indices = host_pool.alloc(fill_len)
self.assertIsNotNone(host_indices, "Host alloc failed")
host_indices = host_indices.to(device="cuda")
host_indices = host_indices.to(device=DEVICE)
self.coordinator.req_to_host_pool[req.kv.req_pool_idx, :fill_len] = host_indices
self.coordinator.req_to_host_pool_allocated_len[req.kv.req_pool_idx] = fill_len
for lid in range(LAYER_NUM):
@@ -273,7 +283,7 @@ class TestHiSparseUnit(unittest.TestCase):
return host_indices
def _build_topk_tokens(self, fill_len, *, include_newest=False):
"""Build a 1-D [TOP_K] int32 cuda tensor of token positions.
"""Build a 1-D [TOP_K] int32 device tensor of token positions.
If include_newest=True, fill_len-1 is guaranteed as the last valid slot.
Pads with -1 when fill_len (or fill_len-1) < TOP_K.
@@ -286,25 +296,25 @@ class TestHiSparseUnit(unittest.TestCase):
"""
n = min(fill_len, TOP_K)
if include_newest and n > 1:
tokens = torch.randperm(fill_len - 1, device="cuda")[: n - 1].to(
tokens = torch.randperm(fill_len - 1, device=DEVICE)[: n - 1].to(
torch.int32
)
tokens = torch.cat(
[tokens, torch.tensor([fill_len - 1], dtype=torch.int32, device="cuda")]
[tokens, torch.tensor([fill_len - 1], dtype=torch.int32, device=DEVICE)]
)
else:
tokens = torch.randperm(fill_len, device="cuda")[:n].to(torch.int32)
tokens = torch.randperm(fill_len, device=DEVICE)[:n].to(torch.int32)
if n < TOP_K:
pad = torch.full((TOP_K - n,), -1, dtype=torch.int32, device="cuda")
pad = torch.full((TOP_K - n,), -1, dtype=torch.int32, device=DEVICE)
tokens = torch.cat([tokens, pad])
return tokens
def _make_batch_tensors(self, reqs, fill_lens):
"""Build (req_pool_indices [int64], seq_lens [int32]) on cuda."""
"""Build (req_pool_indices [int64], seq_lens [int32]) on the active device."""
rpi = torch.tensor(
[r.kv.req_pool_idx for r in reqs], dtype=torch.int64, device="cuda"
[r.kv.req_pool_idx for r in reqs], dtype=torch.int64, device=DEVICE
)
sls = torch.tensor(fill_lens, dtype=torch.int32, device="cuda")
sls = torch.tensor(fill_lens, dtype=torch.int32, device=DEVICE)
return rpi, sls
def _assert_kv_correct(self, locs_row, tokens_row, layer_id, count, msg=""):
@@ -467,7 +477,7 @@ class TestHiSparseUnit(unittest.TestCase):
# Step 1: load the first TOP_K positions from host (no newest token —
# the reserved slot is only valid after map_last_loc_to_buffer which is
# called during an actual decode step, not modelled here).
tokens_s1 = torch.arange(TOP_K, dtype=torch.int32, device="cuda")
tokens_s1 = torch.arange(TOP_K, dtype=torch.int32, device=DEVICE)
locs1 = self._swap_in_selected_pages(
rpi, sls, tokens_s1.unsqueeze(0), layer_id=0
)
@@ -481,7 +491,7 @@ class TestHiSparseUnit(unittest.TestCase):
[
tokens_s1[:half], # hits
torch.arange(
new_start, new_start + half, dtype=torch.int32, device="cuda"
new_start, new_start + half, dtype=torch.int32, device=DEVICE
), # misses
]
)
@@ -633,7 +643,7 @@ class TestHiSparseUnit(unittest.TestCase):
self.coordinator.admit_request_into_staging(req)
self.assertTrue(req.hisparse_staging)
torch.cuda.synchronize()
get_device_module().synchronize()
ready = self.coordinator.collect_ready_reqs()
self.assertEqual(len(ready), 1)
self.assertFalse(req.hisparse_staging)
@@ -669,7 +679,7 @@ class TestHiSparseUnit(unittest.TestCase):
self._write_device_patterns(kv_loc, fill_len)
self.coordinator.admit_request_into_staging(req)
torch.cuda.synchronize()
get_device_module().synchronize()
ready = self.coordinator.collect_ready_reqs()
self.assertEqual(ready, [req])