[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:
Libin Tang
2026-07-24 08:41:01 +08:00
committed by GitHub
co-authored by Christopher Manteuffel Claude Opus 4.8 Ma Mingfei
parent 2f823a2eee
commit 1e10ec93b3
6 changed files with 633 additions and 16 deletions
@@ -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.runtime_context import get_server_args
from sglang.srt.utils import create_device_stream, device_stream_context
try:
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
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__(
self,
num_layers: int,
load_stream: torch.cuda.Stream,
load_stream: torch.Stream,
lmc_connector: LMCacheLayerwiseConnector,
printable: bool = False,
):
@@ -78,7 +79,7 @@ class LayerTransferCounter:
def wait_until(self, layer_id: int):
# Ensure ordering of the async loads wrt compute stream(s).
self.load_stream.synchronize()
with self.load_stream:
with device_stream_context(self.load_stream):
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,
)
self.load_stream = torch.cuda.Stream()
self.store_stream = torch.cuda.Stream()
self.load_stream = create_device_stream(self.device)
self.store_stream = create_device_stream(self.device)
# MP is the default. To use the in-process layerwise connector,
# set ``self._mode = LMCacheMode.IP`` here.
self._mode = LMCacheMode.MP
# MP (multi-process) is the default. XPU defaults to IP (in-process
# layerwise) because the MP connector shares the KV cache via CUDA IPC
# (``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 not cli_lmc_cfg:
raise ValueError(
@@ -351,6 +353,8 @@ class LMCRadixCache(RadixCache):
slot_mapping[:value_numel].fill_(-1)
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)
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
load_stream so the compute stream observes the writes.
"""
self.load_stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(self.load_stream):
current_stream = torch.get_device_module(self.device).current_stream()
self.load_stream.wait_stream(current_stream)
with device_stream_context(self.load_stream):
n = self.lmcache_connector.retrieve_kv(
LoadMetadata(
token_ids=marker.key.token_ids,
@@ -403,7 +408,7 @@ class LMCRadixCache(RadixCache):
request_id=request_id,
)
)
torch.cuda.current_stream().wait_stream(self.load_stream)
current_stream.wait_stream(self.load_stream)
return n
def _ip_load_back(
@@ -419,7 +424,7 @@ class LMCRadixCache(RadixCache):
``start_load_kv`` enqueues the first layer's transfer; the
``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(
LoadMetadata(
token_ids=token_ids,
@@ -472,14 +477,15 @@ class LMCRadixCache(RadixCache):
offset=0,
request_id=req.rid,
)
with torch.cuda.stream(self.store_stream):
self.lmcache_connector.store_kv(store_md)
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.
self._mp_load_back_markers.pop(req.rid, None)
self.dec_lock_ref(new_last_node)
self.lmcache_connector.end_session(req.rid)
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().
with self._node_lock:
self._in_flight_nodes.append(new_last_node)
+12
View File
@@ -545,6 +545,18 @@ def 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():
try:
# Run rocm-smi and capture the output