[XPU] Add XPU device support for LMCache radix cache integration (#23534)
Co-authored-by: Christopher Manteuffel <christopher.manteuffel@intel.com> Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Co-authored-by: Ma Mingfei <mingfei.ma@intel.com>
This commit is contained in:
co-authored by
Christopher Manteuffel
Claude Opus 4.8
Ma Mingfei
parent
2f823a2eee
commit
1e10ec93b3
@@ -103,7 +103,7 @@ jobs:
|
|||||||
timeout-minutes: 60
|
timeout-minutes: 60
|
||||||
run: |
|
run: |
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install --upgrade pip
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install --upgrade pip
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install pytest expecttest ray huggingface_hub tabulate
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install pytest expecttest ray huggingface_hub tabulate "lmcache>=0.3.9"
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip uninstall -y flashinfer-python sgl-kernel sglang
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip uninstall -y flashinfer-python sgl-kernel sglang
|
||||||
docker exec ci_sglang_xpu cp /sglang-checkout/python/pyproject_xpu.toml /sglang-checkout/python/pyproject.toml
|
docker exec ci_sglang_xpu cp /sglang-checkout/python/pyproject_xpu.toml /sglang-checkout/python/pyproject.toml
|
||||||
docker exec -w /sglang-checkout/python ci_sglang_xpu /opt/venv/bin/python3 -m pip install --no-cache-dir . --extra-index-url https://download.pytorch.org/whl/xpu
|
docker exec -w /sglang-checkout/python ci_sglang_xpu /opt/venv/bin/python3 -m pip install --no-cache-dir . --extra-index-url https://download.pytorch.org/whl/xpu
|
||||||
@@ -179,7 +179,7 @@ jobs:
|
|||||||
timeout-minutes: 60
|
timeout-minutes: 60
|
||||||
run: |
|
run: |
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install --upgrade pip
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install --upgrade pip
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install pytest expecttest ray huggingface_hub tabulate
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip install pytest expecttest ray huggingface_hub tabulate "lmcache>=0.3.9"
|
||||||
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip uninstall -y flashinfer-python sgl-kernel sglang
|
docker exec ci_sglang_xpu /opt/venv/bin/python3 -m pip uninstall -y flashinfer-python sgl-kernel sglang
|
||||||
docker exec ci_sglang_xpu cp /sglang-checkout/python/pyproject_xpu.toml /sglang-checkout/python/pyproject.toml
|
docker exec ci_sglang_xpu cp /sglang-checkout/python/pyproject_xpu.toml /sglang-checkout/python/pyproject.toml
|
||||||
docker exec -w /sglang-checkout/python ci_sglang_xpu /opt/venv/bin/python3 -m pip install --no-cache-dir . --extra-index-url https://download.pytorch.org/whl/xpu
|
docker exec -w /sglang-checkout/python ci_sglang_xpu /opt/venv/bin/python3 -m pip install --no-cache-dir . --extra-index-url https://download.pytorch.org/whl/xpu
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
|
from sglang.srt.utils import create_device_stream, device_stream_context
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
|
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
|
||||||
@@ -61,13 +62,13 @@ class LayerTransferCounter:
|
|||||||
|
|
||||||
The KV pool calls `wait_until(layer_id)` after finishing a layer, which we
|
The KV pool calls `wait_until(layer_id)` after finishing a layer, which we
|
||||||
translate into a `load_kv_layerwise(layer_id)` call on the LMCache connector
|
translate into a `load_kv_layerwise(layer_id)` call on the LMCache connector
|
||||||
within the provided CUDA stream.
|
within the provided device stream.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
num_layers: int,
|
num_layers: int,
|
||||||
load_stream: torch.cuda.Stream,
|
load_stream: torch.Stream,
|
||||||
lmc_connector: LMCacheLayerwiseConnector,
|
lmc_connector: LMCacheLayerwiseConnector,
|
||||||
printable: bool = False,
|
printable: bool = False,
|
||||||
):
|
):
|
||||||
@@ -78,7 +79,7 @@ class LayerTransferCounter:
|
|||||||
def wait_until(self, layer_id: int):
|
def wait_until(self, layer_id: int):
|
||||||
# Ensure ordering of the async loads wrt compute stream(s).
|
# Ensure ordering of the async loads wrt compute stream(s).
|
||||||
self.load_stream.synchronize()
|
self.load_stream.synchronize()
|
||||||
with self.load_stream:
|
with device_stream_context(self.load_stream):
|
||||||
self.lmc_connector.load_kv_layerwise(layer_id)
|
self.lmc_connector.load_kv_layerwise(layer_id)
|
||||||
|
|
||||||
|
|
||||||
@@ -131,12 +132,13 @@ class LMCRadixCache(RadixCache):
|
|||||||
tp_group=tp_group.device_group if tp_group is not None else None,
|
tp_group=tp_group.device_group if tp_group is not None else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.load_stream = torch.cuda.Stream()
|
self.load_stream = create_device_stream(self.device)
|
||||||
self.store_stream = torch.cuda.Stream()
|
self.store_stream = create_device_stream(self.device)
|
||||||
|
|
||||||
# MP is the default. To use the in-process layerwise connector,
|
# MP (multi-process) is the default. XPU defaults to IP (in-process
|
||||||
# set ``self._mode = LMCacheMode.IP`` here.
|
# layerwise) because the MP connector shares the KV cache via CUDA IPC
|
||||||
self._mode = LMCacheMode.MP
|
# (``Tensor._share_cuda_``), which is unavailable on XPU.
|
||||||
|
self._mode = LMCacheMode.IP if self.device.type == "xpu" else LMCacheMode.MP
|
||||||
if self._mode is LMCacheMode.MP:
|
if self._mode is LMCacheMode.MP:
|
||||||
if not cli_lmc_cfg:
|
if not cli_lmc_cfg:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -351,6 +353,8 @@ class LMCRadixCache(RadixCache):
|
|||||||
slot_mapping[:value_numel].fill_(-1)
|
slot_mapping[:value_numel].fill_(-1)
|
||||||
slot_mapping[value_numel:].copy_(token_slots)
|
slot_mapping[value_numel:].copy_(token_slots)
|
||||||
|
|
||||||
|
# Dispatch to the mode-specific loader (IP: start_load_kv, MP:
|
||||||
|
# retrieve_kv). Each loader manages its own load_stream context.
|
||||||
num_retrieved = load_fn(slot_mapping, prefix_pad)
|
num_retrieved = load_fn(slot_mapping, prefix_pad)
|
||||||
logger.debug("num_retrieved_tokens: %s", num_retrieved)
|
logger.debug("num_retrieved_tokens: %s", num_retrieved)
|
||||||
|
|
||||||
@@ -392,8 +396,9 @@ class LMCRadixCache(RadixCache):
|
|||||||
"""MP non-layerwise loader: fire ``retrieve_kv`` and wait for the
|
"""MP non-layerwise loader: fire ``retrieve_kv`` and wait for the
|
||||||
load_stream so the compute stream observes the writes.
|
load_stream so the compute stream observes the writes.
|
||||||
"""
|
"""
|
||||||
self.load_stream.wait_stream(torch.cuda.current_stream())
|
current_stream = torch.get_device_module(self.device).current_stream()
|
||||||
with torch.cuda.stream(self.load_stream):
|
self.load_stream.wait_stream(current_stream)
|
||||||
|
with device_stream_context(self.load_stream):
|
||||||
n = self.lmcache_connector.retrieve_kv(
|
n = self.lmcache_connector.retrieve_kv(
|
||||||
LoadMetadata(
|
LoadMetadata(
|
||||||
token_ids=marker.key.token_ids,
|
token_ids=marker.key.token_ids,
|
||||||
@@ -403,7 +408,7 @@ class LMCRadixCache(RadixCache):
|
|||||||
request_id=request_id,
|
request_id=request_id,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
torch.cuda.current_stream().wait_stream(self.load_stream)
|
current_stream.wait_stream(self.load_stream)
|
||||||
return n
|
return n
|
||||||
|
|
||||||
def _ip_load_back(
|
def _ip_load_back(
|
||||||
@@ -419,7 +424,7 @@ class LMCRadixCache(RadixCache):
|
|||||||
``start_load_kv`` enqueues the first layer's transfer; the
|
``start_load_kv`` enqueues the first layer's transfer; the
|
||||||
``LayerTransferCounter`` hook drives the rest during forward.
|
``LayerTransferCounter`` hook drives the rest during forward.
|
||||||
"""
|
"""
|
||||||
with torch.cuda.stream(self.load_stream):
|
with device_stream_context(self.load_stream):
|
||||||
return self.lmcache_connector.start_load_kv(
|
return self.lmcache_connector.start_load_kv(
|
||||||
LoadMetadata(
|
LoadMetadata(
|
||||||
token_ids=token_ids,
|
token_ids=token_ids,
|
||||||
@@ -472,14 +477,15 @@ class LMCRadixCache(RadixCache):
|
|||||||
offset=0,
|
offset=0,
|
||||||
request_id=req.rid,
|
request_id=req.rid,
|
||||||
)
|
)
|
||||||
with torch.cuda.stream(self.store_stream):
|
|
||||||
self.lmcache_connector.store_kv(store_md)
|
|
||||||
if self._mode is LMCacheMode.MP:
|
if self._mode is LMCacheMode.MP:
|
||||||
|
self.lmcache_connector.store_kv(store_md)
|
||||||
# MP store_kv blocks until the daemon's signal event fires, so the slots are safe to evict immediately.
|
# MP store_kv blocks until the daemon's signal event fires, so the slots are safe to evict immediately.
|
||||||
self._mp_load_back_markers.pop(req.rid, None)
|
self._mp_load_back_markers.pop(req.rid, None)
|
||||||
self.dec_lock_ref(new_last_node)
|
self.dec_lock_ref(new_last_node)
|
||||||
self.lmcache_connector.end_session(req.rid)
|
self.lmcache_connector.end_session(req.rid)
|
||||||
elif self._mode is LMCacheMode.IP:
|
elif self._mode is LMCacheMode.IP:
|
||||||
|
with device_stream_context(self.store_stream):
|
||||||
|
self.lmcache_connector.store_kv(store_md)
|
||||||
# Layerwise store is async on store_stream; defer the unlock to evict()'s store_stream.synchronize().
|
# Layerwise store is async on store_stream; defer the unlock to evict()'s store_stream.synchronize().
|
||||||
with self._node_lock:
|
with self._node_lock:
|
||||||
self._in_flight_nodes.append(new_last_node)
|
self._in_flight_nodes.append(new_last_node)
|
||||||
|
|||||||
@@ -545,6 +545,18 @@ def get_device_module():
|
|||||||
return torch.get_device_module()
|
return torch.get_device_module()
|
||||||
|
|
||||||
|
|
||||||
|
def create_device_stream(device):
|
||||||
|
"""Create a device stream for the given device type."""
|
||||||
|
if not isinstance(device, torch.device):
|
||||||
|
device = torch.device(device)
|
||||||
|
return torch.get_device_module(device).Stream(device=device)
|
||||||
|
|
||||||
|
|
||||||
|
def device_stream_context(stream):
|
||||||
|
"""Return the appropriate stream context manager for ``stream``."""
|
||||||
|
return torch.get_device_module(stream.device).stream(stream)
|
||||||
|
|
||||||
|
|
||||||
def get_amdgpu_memory_capacity():
|
def get_amdgpu_memory_capacity():
|
||||||
try:
|
try:
|
||||||
# Run rocm-smi and capture the output
|
# Run rocm-smi and capture the output
|
||||||
|
|||||||
@@ -0,0 +1,361 @@
|
|||||||
|
"""
|
||||||
|
XPU integration tests for LMCache connector in SGLang.
|
||||||
|
|
||||||
|
Tests store/retrieve round-trip on Intel XPU using pure PyTorch ops
|
||||||
|
(index_copy_, index_select) instead of CUDA lmc_ops kernels.
|
||||||
|
|
||||||
|
Uses a single shared connector to avoid LMCacheEngineBuilder singleton
|
||||||
|
issues (close() does not remove from _instances, so re-creating a
|
||||||
|
connector returns a dead engine).
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python3 -m unittest registered.xpu.test_lmcache_connector
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_xpu_ci
|
||||||
|
|
||||||
|
# Must be set before lmcache imports. Save prior values so tearDownModule can
|
||||||
|
# restore them and avoid leaking into other tests in the same process.
|
||||||
|
_PATCHED_ENV = {
|
||||||
|
"LMCACHE_USE_EXPERIMENTAL": "True",
|
||||||
|
"LMCACHE_CONFIG_FILE": os.path.join(
|
||||||
|
os.path.dirname(__file__), "test_lmcache_connector_config.yaml"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
_OLD_ENV = {k: os.environ.get(k) for k in _PATCHED_ENV}
|
||||||
|
os.environ["LMCACHE_USE_EXPERIMENTAL"] = _PATCHED_ENV["LMCACHE_USE_EXPERIMENTAL"]
|
||||||
|
os.environ.setdefault("LMCACHE_CONFIG_FILE", _PATCHED_ENV["LMCACHE_CONFIG_FILE"])
|
||||||
|
|
||||||
|
|
||||||
|
def tearDownModule():
|
||||||
|
for key, old_value in _OLD_ENV.items():
|
||||||
|
if old_value is None:
|
||||||
|
os.environ.pop(key, None)
|
||||||
|
else:
|
||||||
|
os.environ[key] = old_value
|
||||||
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
from lmcache.integration.sglang.sglang_adapter import (
|
||||||
|
LMCacheLayerwiseConnector,
|
||||||
|
LoadMetadata,
|
||||||
|
StoreMetadata,
|
||||||
|
)
|
||||||
|
except ImportError:
|
||||||
|
raise RuntimeError(
|
||||||
|
"LMCache is not installed. "
|
||||||
|
"Install with: NO_CUDA_EXT=1 pip install -e . --no-build-isolation"
|
||||||
|
)
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
|
||||||
|
XPU_AVAILABLE = hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||||
|
|
||||||
|
register_xpu_ci(est_time=60, suite="stage-b-test-1-gpu-xpu")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(XPU_AVAILABLE, "Intel XPU not available")
|
||||||
|
class TestLMCacheXPUConnector(unittest.TestCase):
|
||||||
|
"""Test LMCache layerwise connector store/retrieve on XPU.
|
||||||
|
|
||||||
|
All tests share a single connector instance to avoid the
|
||||||
|
LMCacheEngineBuilder singleton issue where close() does not
|
||||||
|
remove the engine from _instances.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DEVICE = "xpu:0"
|
||||||
|
BUFFER_SIZE = 256
|
||||||
|
INPUT_LEN = 16
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model_config = ModelConfig(model_path="Qwen/Qwen3-4B")
|
||||||
|
cls.head_num = cls.model_config.num_key_value_heads
|
||||||
|
cls.head_dim = cls.model_config.head_dim
|
||||||
|
cls.layer_num = cls.model_config.num_hidden_layers
|
||||||
|
cls.vocab_size = cls.model_config.vocab_size
|
||||||
|
|
||||||
|
# Shared KV buffers and connector (created once)
|
||||||
|
cls.k_buffer = [
|
||||||
|
torch.randn(
|
||||||
|
cls.BUFFER_SIZE,
|
||||||
|
cls.head_num,
|
||||||
|
cls.head_dim,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=cls.DEVICE,
|
||||||
|
)
|
||||||
|
for _ in range(cls.layer_num)
|
||||||
|
]
|
||||||
|
cls.v_buffer = [
|
||||||
|
torch.randn(
|
||||||
|
cls.BUFFER_SIZE,
|
||||||
|
cls.head_num,
|
||||||
|
cls.head_dim,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=cls.DEVICE,
|
||||||
|
)
|
||||||
|
for _ in range(cls.layer_num)
|
||||||
|
]
|
||||||
|
cls.connector = LMCacheLayerwiseConnector(
|
||||||
|
cls.model_config,
|
||||||
|
tp_size=1,
|
||||||
|
rank=0,
|
||||||
|
k_pool=cls.k_buffer,
|
||||||
|
v_pool=cls.v_buffer,
|
||||||
|
config_file=os.environ["LMCACHE_CONFIG_FILE"],
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
cls.connector.close()
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Re-randomize buffers before each test for isolation."""
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.k_buffer[i].normal_()
|
||||||
|
self.v_buffer[i].normal_()
|
||||||
|
|
||||||
|
def _unique_tokens(self, length=None, salt=0):
|
||||||
|
"""Generate unique token ids unlikely to collide across tests."""
|
||||||
|
n = length or self.INPUT_LEN
|
||||||
|
base = torch.randint(0, self.vocab_size, (n,))
|
||||||
|
return [(t.item() + salt) % self.vocab_size for t in base]
|
||||||
|
|
||||||
|
def test_store_then_retrieve(self):
|
||||||
|
"""Basic: store KV, clear buffers, retrieve and verify match."""
|
||||||
|
token_ids = self._unique_tokens(salt=100)
|
||||||
|
kv_indices = torch.randint(0, self.BUFFER_SIZE, (self.INPUT_LEN,))
|
||||||
|
|
||||||
|
# First retrieve should return 0 (cold cache)
|
||||||
|
load_meta = LoadMetadata(
|
||||||
|
token_ids=token_ids,
|
||||||
|
slot_mapping=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.assertEqual(self.connector.start_load_kv(load_meta), 0)
|
||||||
|
|
||||||
|
# Store
|
||||||
|
store_meta = StoreMetadata(
|
||||||
|
last_node=None,
|
||||||
|
token_ids=token_ids,
|
||||||
|
kv_indices=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.connector.store_kv(store_meta)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
# Save ground truth before clearing
|
||||||
|
gt_k = [self.k_buffer[i][kv_indices].clone() for i in range(self.layer_num)]
|
||||||
|
gt_v = [self.v_buffer[i][kv_indices].clone() for i in range(self.layer_num)]
|
||||||
|
|
||||||
|
# Clear buffers
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.k_buffer[i].zero_()
|
||||||
|
self.v_buffer[i].zero_()
|
||||||
|
|
||||||
|
# Retrieve
|
||||||
|
ret = self.connector.start_load_kv(load_meta)
|
||||||
|
self.assertEqual(ret, self.INPUT_LEN)
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
self.connector.load_kv_layerwise(i)
|
||||||
|
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
# Verify
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
actual_k = self.k_buffer[i][kv_indices]
|
||||||
|
actual_v = self.v_buffer[i][kv_indices]
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_k, gt_k[i]),
|
||||||
|
f"Layer {i}: K mismatch (max diff {(actual_k - gt_k[i]).abs().max():.6f})",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_v, gt_v[i]),
|
||||||
|
f"Layer {i}: V mismatch (max diff {(actual_v - gt_v[i]).abs().max():.6f})",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_retrieve_cold_cache_returns_zero(self):
|
||||||
|
"""Retrieve from empty cache should return 0 tokens."""
|
||||||
|
token_ids = self._unique_tokens(salt=200)
|
||||||
|
kv_indices = torch.randint(0, self.BUFFER_SIZE, (self.INPUT_LEN,))
|
||||||
|
|
||||||
|
load_meta = LoadMetadata(
|
||||||
|
token_ids=token_ids,
|
||||||
|
slot_mapping=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.assertEqual(self.connector.start_load_kv(load_meta), 0)
|
||||||
|
|
||||||
|
def test_slot_mapping_with_negative_indices(self):
|
||||||
|
"""slot_mapping may contain -1 for an already-cached prefix.
|
||||||
|
|
||||||
|
The -1 filtering fix must (a) not crash on XPU and (b) leave the -1
|
||||||
|
slots untouched while correctly restoring the valid slots.
|
||||||
|
"""
|
||||||
|
# Distinct index for every token so -1 slots and valid slots never
|
||||||
|
# alias each other, keeping the untouched/restored assertions exact.
|
||||||
|
kv_indices = torch.randperm(self.BUFFER_SIZE)[: self.INPUT_LEN]
|
||||||
|
token_ids = self._unique_tokens(salt=400)
|
||||||
|
|
||||||
|
# Store first
|
||||||
|
store_meta = StoreMetadata(
|
||||||
|
last_node=None,
|
||||||
|
token_ids=token_ids,
|
||||||
|
kv_indices=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.connector.store_kv(store_meta)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
# Build slot_mapping with -1 prefix (simulating already-cached tokens).
|
||||||
|
num_cached = 8
|
||||||
|
slot_mapping_with_neg = kv_indices.clone()
|
||||||
|
slot_mapping_with_neg[:num_cached] = -1
|
||||||
|
cached_slots = kv_indices[:num_cached]
|
||||||
|
valid_slots = kv_indices[num_cached:]
|
||||||
|
|
||||||
|
# Ground truth for the valid tail (what retrieve must restore).
|
||||||
|
gt_k = [self.k_buffer[i][valid_slots].clone() for i in range(self.layer_num)]
|
||||||
|
gt_v = [self.v_buffer[i][valid_slots].clone() for i in range(self.layer_num)]
|
||||||
|
|
||||||
|
# Clear buffers: untouched -1 slots stay zero, valid slots get restored.
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.k_buffer[i].zero_()
|
||||||
|
self.v_buffer[i].zero_()
|
||||||
|
|
||||||
|
load_meta = LoadMetadata(
|
||||||
|
token_ids=token_ids,
|
||||||
|
slot_mapping=slot_mapping_with_neg,
|
||||||
|
# offset marks how many leading tokens are already cached (the
|
||||||
|
# -1 prefix in slot_mapping); this mirrors how lmc_radix_cache's
|
||||||
|
# _ip_load_back derives offset from the already-matched prefix.
|
||||||
|
offset=num_cached,
|
||||||
|
)
|
||||||
|
# Should not crash on XPU (the -1 filtering fix is critical here).
|
||||||
|
ret = self.connector.start_load_kv(load_meta)
|
||||||
|
self.assertEqual(ret, self.INPUT_LEN - num_cached)
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
self.connector.load_kv_layerwise(i)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
# -1 positions must be untouched (still zero).
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(self.k_buffer[i][cached_slots] == 0),
|
||||||
|
f"Layer {i}: -1 K slots were written",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.all(self.v_buffer[i][cached_slots] == 0),
|
||||||
|
f"Layer {i}: -1 V slots were written",
|
||||||
|
)
|
||||||
|
# Valid positions must be restored.
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(self.k_buffer[i][valid_slots], gt_k[i]),
|
||||||
|
f"Layer {i}: valid K slots not restored",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(self.v_buffer[i][valid_slots], gt_v[i]),
|
||||||
|
f"Layer {i}: valid V slots not restored",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_multiple_store_retrieve_cycles(self):
|
||||||
|
"""Multiple store/retrieve cycles should not leak or corrupt."""
|
||||||
|
for cycle in range(3):
|
||||||
|
token_ids = self._unique_tokens(salt=500 + cycle * 1000)
|
||||||
|
kv_indices = torch.randint(0, self.BUFFER_SIZE, (self.INPUT_LEN,))
|
||||||
|
|
||||||
|
store_meta = StoreMetadata(
|
||||||
|
last_node=None,
|
||||||
|
token_ids=token_ids,
|
||||||
|
kv_indices=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.connector.store_kv(store_meta)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
gt_k = [self.k_buffer[i][kv_indices].clone() for i in range(self.layer_num)]
|
||||||
|
gt_v = [self.v_buffer[i][kv_indices].clone() for i in range(self.layer_num)]
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.k_buffer[i].zero_()
|
||||||
|
self.v_buffer[i].zero_()
|
||||||
|
|
||||||
|
load_meta = LoadMetadata(
|
||||||
|
token_ids=token_ids,
|
||||||
|
slot_mapping=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
ret = self.connector.start_load_kv(load_meta)
|
||||||
|
self.assertEqual(
|
||||||
|
ret,
|
||||||
|
self.INPUT_LEN,
|
||||||
|
f"Cycle {cycle}: expected {self.INPUT_LEN}, got {ret}",
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
self.connector.load_kv_layerwise(i)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
actual_k = self.k_buffer[i][kv_indices]
|
||||||
|
actual_v = self.v_buffer[i][kv_indices]
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_k, gt_k[i]),
|
||||||
|
f"Cycle {cycle}, Layer {i}: K mismatch",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_v, gt_v[i]),
|
||||||
|
f"Cycle {cycle}, Layer {i}: V mismatch",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_bf16_dtype_preserved(self):
|
||||||
|
"""Verify bf16 dtype is preserved through store->retrieve cycle.
|
||||||
|
|
||||||
|
TODO: once XPU LMCache supports KV quantization, extend the dtype
|
||||||
|
coverage here to the quantized store/retrieve dtypes (e.g. fp8).
|
||||||
|
"""
|
||||||
|
token_ids = self._unique_tokens(salt=600)
|
||||||
|
kv_indices = torch.randint(0, self.BUFFER_SIZE, (self.INPUT_LEN,))
|
||||||
|
|
||||||
|
store_meta = StoreMetadata(
|
||||||
|
last_node=None,
|
||||||
|
token_ids=token_ids,
|
||||||
|
kv_indices=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
self.connector.store_kv(store_meta)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.k_buffer[i].zero_()
|
||||||
|
self.v_buffer[i].zero_()
|
||||||
|
|
||||||
|
load_meta = LoadMetadata(
|
||||||
|
token_ids=token_ids,
|
||||||
|
slot_mapping=kv_indices,
|
||||||
|
offset=0,
|
||||||
|
)
|
||||||
|
ret = self.connector.start_load_kv(load_meta)
|
||||||
|
self.assertEqual(ret, self.INPUT_LEN)
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
self.connector.load_kv_layerwise(i)
|
||||||
|
torch.xpu.synchronize()
|
||||||
|
|
||||||
|
for i in range(self.layer_num):
|
||||||
|
self.assertEqual(self.k_buffer[i].dtype, torch.bfloat16)
|
||||||
|
self.assertEqual(self.v_buffer[i].dtype, torch.bfloat16)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
# LMCache config for XPU unit tests
|
||||||
|
# save_unfull_chunk enables storing sequences shorter than chunk_size
|
||||||
|
chunk_size: 8
|
||||||
|
local_cpu: true
|
||||||
|
use_layerwise: true
|
||||||
|
max_local_cpu_size: 1
|
||||||
|
save_unfull_chunk: true
|
||||||
@@ -0,0 +1,231 @@
|
|||||||
|
"""
|
||||||
|
XPU integration tests for LMCRadixCache (IP mode).
|
||||||
|
|
||||||
|
Unlike test_lmcache_connector.py, which drives LMCacheLayerwiseConnector
|
||||||
|
directly, this test builds a real ReqToTokenPool + MHATokenToKVPool +
|
||||||
|
TokenToKVPoolAllocator and drives LMCRadixCache itself through the request
|
||||||
|
lifecycle (match_prefix -> cache_finished_req, then evict a fresh cache to
|
||||||
|
force a real LMCache load-back through match_prefix again). This exercises
|
||||||
|
LayerTransferCounter.wait_until, _load_back's slot/offset math, and the
|
||||||
|
store/load stream synchronization that the connector-level tests never touch.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python3 -m unittest registered.xpu.test_lmcache_radix_cache
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_xpu_ci
|
||||||
|
|
||||||
|
# Must be set before lmcache imports. Save prior values so tearDownModule can
|
||||||
|
# restore them and avoid leaking into other tests in the same process.
|
||||||
|
_PATCHED_ENV = {
|
||||||
|
"LMCACHE_USE_EXPERIMENTAL": "True",
|
||||||
|
"LMCACHE_CONFIG_FILE": os.path.join(
|
||||||
|
os.path.dirname(__file__), "test_lmcache_connector_config.yaml"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
_OLD_ENV = {k: os.environ.get(k) for k in _PATCHED_ENV}
|
||||||
|
os.environ["LMCACHE_USE_EXPERIMENTAL"] = _PATCHED_ENV["LMCACHE_USE_EXPERIMENTAL"]
|
||||||
|
os.environ.setdefault("LMCACHE_CONFIG_FILE", _PATCHED_ENV["LMCACHE_CONFIG_FILE"])
|
||||||
|
|
||||||
|
|
||||||
|
def tearDownModule():
|
||||||
|
for key, old_value in _OLD_ENV.items():
|
||||||
|
if old_value is None:
|
||||||
|
os.environ.pop(key, None)
|
||||||
|
else:
|
||||||
|
os.environ[key] = old_value
|
||||||
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
import lmcache.integration.sglang.sglang_adapter # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
raise RuntimeError("LMCache is not installed. Install with: pip install lmcache")
|
||||||
|
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator
|
||||||
|
from sglang.srt.mem_cache.base_prefix_cache import EvictParams, MatchPrefixParams
|
||||||
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
|
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
|
||||||
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
|
from sglang.srt.mem_cache.storage.lmcache.lmc_radix_cache import LMCRadixCache
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
|
XPU_AVAILABLE = hasattr(torch, "xpu") and torch.xpu.is_available()
|
||||||
|
|
||||||
|
register_xpu_ci(est_time=60, suite="stage-b-test-1-gpu-xpu")
|
||||||
|
|
||||||
|
|
||||||
|
def _make_req(rid, req_pool_idx, token_ids, tree):
|
||||||
|
"""Fake Req with the fields LMCRadixCache/RadixCache read (mirrors the
|
||||||
|
SimpleNamespace pattern in test_swa_eviction_boundary.py)."""
|
||||||
|
req = SimpleNamespace(
|
||||||
|
rid=rid,
|
||||||
|
req_pool_idx=req_pool_idx,
|
||||||
|
origin_input_ids=token_ids,
|
||||||
|
output_ids=[],
|
||||||
|
extra_key=None,
|
||||||
|
last_node=tree.root_node,
|
||||||
|
cache_protected_len=0,
|
||||||
|
priority=0,
|
||||||
|
kv_committed_freed=False,
|
||||||
|
kv_committed_len=len(token_ids),
|
||||||
|
)
|
||||||
|
req.pop_committed_kv_cache = lambda: len(token_ids)
|
||||||
|
return req
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(XPU_AVAILABLE, "Intel XPU not available")
|
||||||
|
class TestLMCRadixCacheXPU(unittest.TestCase):
|
||||||
|
"""Drive LMCRadixCache (IP mode) through match_prefix/cache_finished_req
|
||||||
|
with a real KV pool, to cover the code path test_lmcache_connector.py
|
||||||
|
(which talks to the connector directly) never exercises."""
|
||||||
|
|
||||||
|
DEVICE = "xpu:0"
|
||||||
|
BUFFER_SIZE = 256
|
||||||
|
MAX_CONTEXT_LEN = 64
|
||||||
|
INPUT_LEN = 16
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model_config = ModelConfig(model_path="Qwen/Qwen3-4B")
|
||||||
|
cls._override = get_context().override_server_args(
|
||||||
|
lmcache_config_file=os.environ["LMCACHE_CONFIG_FILE"],
|
||||||
|
speculative_eagle_topk=None,
|
||||||
|
)
|
||||||
|
cls._override.install()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
cls._override.restore()
|
||||||
|
|
||||||
|
def _build_tree(self):
|
||||||
|
model_config = self.model_config
|
||||||
|
kv_pool = MHATokenToKVPool(
|
||||||
|
size=self.BUFFER_SIZE,
|
||||||
|
page_size=1,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
head_num=model_config.num_key_value_heads,
|
||||||
|
head_dim=model_config.head_dim,
|
||||||
|
layer_num=model_config.num_hidden_layers,
|
||||||
|
device=self.DEVICE,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
)
|
||||||
|
allocator = TokenToKVPoolAllocator(
|
||||||
|
size=self.BUFFER_SIZE,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=self.DEVICE,
|
||||||
|
kvcache=kv_pool,
|
||||||
|
need_sort=False,
|
||||||
|
)
|
||||||
|
req_to_token_pool = ReqToTokenPool(
|
||||||
|
size=8,
|
||||||
|
max_context_len=self.MAX_CONTEXT_LEN,
|
||||||
|
device=self.DEVICE,
|
||||||
|
enable_memory_saver=False,
|
||||||
|
)
|
||||||
|
tree = LMCRadixCache(
|
||||||
|
params=CacheInitParams(
|
||||||
|
disable=False,
|
||||||
|
req_to_token_pool=req_to_token_pool,
|
||||||
|
token_to_kv_pool_allocator=allocator,
|
||||||
|
page_size=1,
|
||||||
|
),
|
||||||
|
model_config=model_config,
|
||||||
|
tp_size=1,
|
||||||
|
rank=0,
|
||||||
|
)
|
||||||
|
return tree, allocator, req_to_token_pool, kv_pool
|
||||||
|
|
||||||
|
def test_store_then_load_back_through_match_prefix(self):
|
||||||
|
"""Full lifecycle: match_prefix (miss) -> cache_finished_req (store to
|
||||||
|
LMCache + insert into radix) -> evict the radix entry -> match_prefix
|
||||||
|
again must retrieve from LMCache via LayerTransferCounter and restore
|
||||||
|
the original KV content, proving _load_back's offset/slot math and
|
||||||
|
the load_stream synchronization are correct end-to-end."""
|
||||||
|
tree, allocator, req_to_token_pool, kv_pool = self._build_tree()
|
||||||
|
try:
|
||||||
|
token_ids = torch.randint(
|
||||||
|
0, self.model_config.vocab_size, (self.INPUT_LEN,)
|
||||||
|
).tolist()
|
||||||
|
|
||||||
|
# No prior entries: match_prefix should be a full miss.
|
||||||
|
miss_res = tree.match_prefix(MatchPrefixParams(key=RadixKey(token_ids)))
|
||||||
|
self.assertEqual(miss_res.device_indices.numel(), 0)
|
||||||
|
|
||||||
|
# Allocate KV slots for the request, write ground-truth K/V, then
|
||||||
|
# commit it as a finished request (inserts into radix + stores to
|
||||||
|
# LMCache on tree.store_stream).
|
||||||
|
req_pool_idx = req_to_token_pool.alloc(
|
||||||
|
[
|
||||||
|
SimpleNamespace(
|
||||||
|
req_pool_idx=None, inflight_middle_chunks=0, kv_committed_len=0
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)[0]
|
||||||
|
kv_slots = allocator.alloc(self.INPUT_LEN)
|
||||||
|
self.assertIsNotNone(kv_slots)
|
||||||
|
req_to_token_pool.write((req_pool_idx, slice(0, self.INPUT_LEN)), kv_slots)
|
||||||
|
|
||||||
|
gt_k = []
|
||||||
|
gt_v = []
|
||||||
|
for layer_id in range(self.model_config.num_hidden_layers):
|
||||||
|
k = torch.randn(
|
||||||
|
self.INPUT_LEN,
|
||||||
|
self.model_config.num_key_value_heads,
|
||||||
|
self.model_config.head_dim,
|
||||||
|
dtype=torch.bfloat16,
|
||||||
|
device=self.DEVICE,
|
||||||
|
)
|
||||||
|
v = torch.randn_like(k)
|
||||||
|
kv_pool.k_buffer[layer_id][kv_slots] = k
|
||||||
|
kv_pool.v_buffer[layer_id][kv_slots] = v
|
||||||
|
gt_k.append(k.clone())
|
||||||
|
gt_v.append(v.clone())
|
||||||
|
|
||||||
|
req = _make_req("req-0", req_pool_idx, token_ids, tree)
|
||||||
|
tree.cache_finished_req(req, kv_len_to_handle=len(token_ids))
|
||||||
|
# IP-mode store is async on tree.store_stream; evict()'s
|
||||||
|
# synchronize() is what the real scheduler relies on to make the
|
||||||
|
# store visible before slots are reused.
|
||||||
|
tree.evict(EvictParams(num_tokens=0))
|
||||||
|
|
||||||
|
# Evict everything from the radix tree so the only remaining copy
|
||||||
|
# of this KV is inside LMCache, forcing a real load-back on the
|
||||||
|
# next match_prefix.
|
||||||
|
tree.evict(EvictParams(num_tokens=self.INPUT_LEN))
|
||||||
|
self.assertEqual(tree.total_size(), 0)
|
||||||
|
|
||||||
|
for layer_id in range(self.model_config.num_hidden_layers):
|
||||||
|
kv_pool.k_buffer[layer_id].zero_()
|
||||||
|
kv_pool.v_buffer[layer_id].zero_()
|
||||||
|
|
||||||
|
reload_res = tree.match_prefix(MatchPrefixParams(key=RadixKey(token_ids)))
|
||||||
|
self.assertEqual(reload_res.device_indices.numel(), self.INPUT_LEN)
|
||||||
|
|
||||||
|
new_slots = reload_res.device_indices
|
||||||
|
for layer_id in range(self.model_config.num_hidden_layers):
|
||||||
|
# get_key_buffer/get_value_buffer (not the raw k_buffer/v_buffer
|
||||||
|
# list) is what invokes layer_transfer_counter.wait_until —
|
||||||
|
# the real per-layer forward hook this test exists to cover.
|
||||||
|
actual_k = kv_pool.get_key_buffer(layer_id)[new_slots]
|
||||||
|
actual_v = kv_pool.get_value_buffer(layer_id)[new_slots]
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_k, gt_k[layer_id]),
|
||||||
|
f"Layer {layer_id}: K not restored via LMCache load-back",
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
torch.allclose(actual_v, gt_v[layer_id]),
|
||||||
|
f"Layer {layer_id}: V not restored via LMCache load-back",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
tree.lmcache_connector.close()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user