Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
34 commits
Select commit Hold shift + click to select a range
f806ad5
Fix unified HiCache physical transfers
Sep 2, 2026
329fb0c
Merge latest unified host-pool stack
Sep 2, 2026
f16293b
Support unified HiCache buffer staging
Sep 4, 2026
54d3892
Merge compacting unified host pool
Sep 4, 2026
136669c
Enable shared unified HiCache staging safely
Sep 4, 2026
63abe66
Merge unified decode host safety gates
Sep 4, 2026
f0af6be
Guard unified HiCache compatibility
Sep 4, 2026
1f8c4fe
Merge review fixes into HiCache transfers
Sep 4, 2026
d4462ed
Merge commit '372d2111cc' into yonghao/ump-hicache-physical-transfers
Sep 4, 2026
f417938
Merge commit 'ec0e7777f38d651b0c3e849d033b81985fc091b3' into yonghao/…
Sep 4, 2026
3c26277
Merge commit 'b0018def8d' into yonghao/ump-hicache-physical-transfers
Sep 4, 2026
18142d6
Keep physical reservation cleanup host-sync-free
Sep 4, 2026
95e9b6d
Merge host-pool updates and preserve unified transfer lifetimes
Sep 8, 2026
097e185
Merge UMP cleanup and restore host-sync-free physical release
Sep 8, 2026
b20d10f
Merge host pool refresh and adapt HiCache physical transfers
Sep 10, 2026
87df0bc
Merge Python 3.10 fixture compatibility fix from host pool
Sep 10, 2026
709e1df
Merge host-pool cleanup into physical transfers
Sep 10, 2026
3870148
Fix SWA-only backup and shared-host prefetch fallback
Sep 12, 2026
b03bcd2
Merge host-pool shared-capacity cleanup
Sep 13, 2026
dc84a3b
Merge shared host transfer cleanup
Sep 13, 2026
320d041
Separate storage hit allocation from queue policy
Sep 13, 2026
809821e
Keep short-prefix threshold policy at the cache boundary
Sep 13, 2026
242578b
Merge shared SWA budget cleanup
Sep 13, 2026
e951bb1
Merge host-pool fixes and preserve eviction progress under host pressure
Sep 13, 2026
1fcd125
Merge UMP cleanup and preserve non-unified HiCache failure semantics
Sep 13, 2026
a9a744a
Merge unified prefill rounding fix into HiCache transfers
Sep 14, 2026
e843a0a
Consolidate UMBP page layout and SWA host-pool dispatch
Sep 14, 2026
8e0df01
Merge host-pool cleanup preserving HiCache bindings
Sep 14, 2026
74b7f42
Merge updated unified host pools and preserve HiCache transfer fixes
Sep 14, 2026
acce638
Merge host-pool updates and simplify HiCache allocation policies
Sep 14, 2026
ccd3487
Fix Rust SWA backup sources and gate sidecar error recovery
Sep 14, 2026
d9c9565
Merge host-pool capability cleanup into HiCache transfers
Sep 14, 2026
a640ec3
Align HiCache prefix reuse with SWA tail allocation
Sep 14, 2026
bbbebd7
Merge shared KV budget fixes into HiCache physical transfers
Sep 14, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 42 additions & 10 deletions python/sglang/srt/arg_groups/kv_cache_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -262,10 +262,6 @@ def handle_unified_memory_pool(server_args: Any) -> None:
)
if cfg.disaggregation_decode_retraction_backup == "host_pool":
model_config = model_config_of(server_args)
assert not cfg.disaggregation_decode_enable_radix_cache, (
"--enable-unified-memory host-pool decode retraction does not "
"support decode radix-cache H2D/D2H transfers yet."
)
assert mambaish_config(model_config) is None, (
"--enable-unified-memory host-pool decode retraction does not "
"support hybrid-Mamba models."
Expand Down Expand Up @@ -305,13 +301,49 @@ def handle_unified_memory_pool(server_args: Any) -> None:
"write loc, so a captured decode replay raises. "
"TODO(ch-wan): carry out_cache_loc_virtual into the child view."
)
assert not (cfg.enable_hierarchical_cache or cfg.enable_lmcache), (
"--enable-unified-memory is not yet compatible with hierarchical / "
"host-tiered KV cache (--enable-hierarchical-cache / --enable-lmcache): "
"the unified-memory-pool init wires up no host pools, and its device mamba / "
"full-attention slots are VIRTUAL — the host-offload path does not "
"translate them to physical."
assert not cfg.enable_lmcache, (
"--enable-unified-memory is not yet compatible with --enable-lmcache."
)
assert not cfg.enable_unified_cache_external_linker, (
"--enable-unified-memory does not support "
"--enable-unified-cache-external-linker: direct L3 transfers do not "
"preserve unified page-envelope indices and compaction lifetimes. "
"Use --enable-hierarchical-cache for supported L2/L3 transfers."
)
if cfg.enable_hierarchical_cache:
model_config = model_config_of(server_args)
assert not use_mla_backend(server_args), (
"--enable-unified-memory with hierarchical cache does not support "
"MLA models yet."
)
assert mambaish_config(model_config) is None, (
"--enable-unified-memory with hierarchical cache does not support "
"recurrent-state models yet."
)
assert cfg.speculative_algorithm is None, (
"--enable-unified-memory with hierarchical cache does not support "
"speculative decoding yet."
)
assert cfg.hicache_io_backend in {"kernel", "direct"}, (
"--enable-unified-memory with hierarchical cache supports only "
"the kernel and direct I/O backends."
)
assert cfg.pp_size == 1, (
"--enable-unified-memory with hierarchical cache does not support "
"pipeline parallelism (--pp-size > 1)."
)
supported_storage_backends = {None, "file", "sim", "mori", "shm"}
if cfg.hicache_storage_backend not in supported_storage_backends:
raise ValueError(
"--enable-unified-memory with hierarchical cache does not "
"support this storage backend yet; supported backends are "
"file, sim, mori, and shm. Got "
f"{cfg.hicache_storage_backend!r}."
)
assert not envs.SGLANG_DISABLE_LAZY_COMPACTION.get(), (
"--enable-unified-memory with hierarchical cache requires lazy "
"compaction so pending H2D physical reservations remain stable."
)
if cfg.dcp_size > 1:
_validate_unified_memory_dcp(server_args)
# Prefill cuda-graph capture IS wired for the unified pool: the captured
Expand Down
29 changes: 12 additions & 17 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -762,20 +762,29 @@ def _match_prefix_and_lock(self, req: Req) -> DecodePrefixMatch:
Match a request against the decode-side radix cache, lock the matched
node to prevent eviction, and return the matched prefix information.
"""
max_prefix_len = None
if self._uses_swa_tail_prealloc():
fill_len = self._pre_alloc_fill_len(req)
max_prefix_len = fill_len - self._swa_tail_len(fill_len)
# Match and lock only reusable FULL KV. The entire SWA tail must be
# freshly allocated, including when the prefix comes from L2/L3.
result = match_prefix_for_req(
self.tree_cache,
req,
req.origin_input_ids,
cow_mamba=self.tree_cache.supports_mamba(),
include_req=True,
max_prefix_len=max_prefix_len,
)
# Keep aggregated scheduling semantics while preserving the SWA lock
# boundary needed for the matching dec_lock_ref; the full receipt
# travels on the req so every later release mirrors this acquire.
req.lock_receipt = self.tree_cache.inc_lock_ref(
result.last_device_node
).to_dec_params()
return self._build_decode_prefix_match(req, result)
return self._build_decode_prefix_match(
req, result, max_prefix_len=max_prefix_len
)

def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]:
prefill_info = self.kv_manager.prefill_info_table.get(_bootstrap_addr(req))
Expand Down Expand Up @@ -1338,20 +1347,6 @@ def pop_preallocated(

fill_len = self._pre_alloc_fill_len(decode_req.req)

# Cap full-attention prefix reuse at the sliding-window start so
# the SWA window lands entirely in the fresh delta, keeping
# alloc_extend_swa_tail's tail->full mapping in range. Costs reuse
# of only the last ~window_size full-attention tokens.
if uses_swa_tail_prealloc and prefix_len > 0:
swa_prefix_cap = fill_len - self._swa_tail_len(fill_len)
if prefix_len > swa_prefix_cap:
prefix_len = swa_prefix_cap
prefix_indices = prefix_indices[:prefix_len]
# Cap the prefill-committed prefix too: tokens past the
# cap are not device-resident, so prefill must transfer
# them.
total_prefix_len = prefix_len

# Decode transfers the SWA tail fresh, so retain only the
# full-attention prefix lock needed for reuse.
if (
Expand Down Expand Up @@ -2127,9 +2122,9 @@ def alloc_for_decode_prealloc(
# the live window tail.
kv_loc = allocator.alloc_extend_swa_tail(
prefix_lens=torch.tensor(
[prefix_len], dtype=torch.int64, device=device
[total_prefix_len], dtype=torch.int64, device=device
),
prefix_lens_cpu=torch.tensor([prefix_len], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([total_prefix_len], dtype=torch.int64),
seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device),
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
last_loc=last_loc,
Expand Down
7 changes: 5 additions & 2 deletions python/sglang/srt/disaggregation/decode_hicache_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,9 @@ class HiCacheRestoreResult(Enum):
class DecodeHiCachePreallocMixin:
"""HiCache hooks for ``DecodePreallocQueue``: issue prefetch + reserve tokens."""

def _build_decode_prefix_match(self, req: Req, result: Any) -> DecodePrefixMatch:
def _build_decode_prefix_match(
self, req: Req, result: Any, *, max_prefix_len: Optional[int] = None
) -> DecodePrefixMatch:
"""Convert a ``match_prefix_for_req`` result into ``DecodePrefixMatch``.

Performs the optional L3 storage hit length query when decode-side
Expand All @@ -78,7 +80,7 @@ def _build_decode_prefix_match(self, req: Req, result: Any) -> DecodePrefixMatch
last_host_node
):
matched_len = l1_prefix_len + l2_host_hit_length
suffix_tokens = req.origin_input_ids[matched_len:]
suffix_tokens = req.origin_input_ids[matched_len:max_prefix_len]
last_hash = self.tree_cache.get_last_hash_value(last_host_node)
prefix_keys = (
self.tree_cache.get_prefix_hash_values(last_host_node)
Expand Down Expand Up @@ -218,6 +220,7 @@ def _try_hicache_queue_load_back(self, dr: DecodeRequest) -> bool:
dr.req.origin_input_ids,
cow_mamba=False,
include_req=True,
max_prefix_len=pm.decode_prefix_len,
)
new_indices, restored_node = self.tree_cache.init_load_back(
InitLoadBackParams(
Expand Down
85 changes: 68 additions & 17 deletions python/sglang/srt/managers/cache_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@
from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
from sglang.srt.mem_cache.utils import get_storage_hash_str
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import get_memory, get_parallel
from sglang.srt.utils import get_device_module

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -822,6 +822,9 @@ def start_writing(self) -> None:
completion = self.l2_transfer_engine.submit_device_to_host(
self._l2_transfers(host_indices, device_indices, pool_transfers)
)
self.mem_pool_device_allocator.set_hicache_transfer_done_event(
(id(self), "write"), completion.finish_event
)

self.ack_write_queue.append(
HiCacheAck(
Expand Down Expand Up @@ -917,6 +920,11 @@ def _move_op_indices(
) -> tuple[torch.Tensor, torch.Tensor, Optional[List[PoolTransfer]]]:
return (*self.move_indices(op.host_indices, op.device_indices), None)

def _move_load_operation(
self, op: CacheOperation
) -> tuple[torch.Tensor, torch.Tensor, Optional[List[PoolTransfer]]]:
return self._move_op_indices(op)

def _l2_transfers(
self,
host_indices: torch.Tensor,
Expand Down Expand Up @@ -947,7 +955,7 @@ def start_loading(self) -> int:

producer_id = self.layer_done_counter.update_producer()
op = CacheOperation.merge_ops(self.load_queue)
host_indices, device_indices, pool_transfers = self._move_op_indices(op)
host_indices, device_indices, pool_transfers = self._move_load_operation(op)
self.load_queue.clear()
producer_event = self.layer_done_counter.events[producer_id]
producer_event.start_event.record()
Expand All @@ -966,6 +974,9 @@ def start_loading(self) -> int:
on_layer_done=producer_event.complete,
layer_num=self.layer_num,
)
self.mem_pool_device_allocator.set_hicache_transfer_done_event(
(id(self), "load"), completion.finish_event
)

self.ack_load_queue.append(
HiCacheAck(
Expand Down Expand Up @@ -1093,13 +1104,22 @@ def _page_transfer(self, operation: PrefetchOperation) -> int:
# Get one batch token, and update the completed_tokens if succeed
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)

hit_pages = self._page_transfer_kv_batch(
operation,
batch_hashes,
batch_host_indices,
extra_info,
kv_derived_transfers,
)
try:
hit_pages = self._page_transfer_kv_batch(
operation,
batch_hashes,
batch_host_indices,
extra_info,
kv_derived_transfers,
)
except Exception:
if not get_memory().enable_unified_memory:
raise
logger.exception(
"HiCache prefetch transfer failed for request %s",
Comment thread
ZYHowell marked this conversation as resolved.
operation.request_id,
)
hit_pages = 0
# Check termination
if hit_pages != len(batch_hashes):
all_success = False
Expand Down Expand Up @@ -1160,19 +1180,20 @@ def prefetch_io_aux_func(self):
while not self.storage_stop_event.is_set():
try:
operation = self.prefetch_buffer.get(block=True, timeout=1)
if operation is None:
continue
except Empty:
continue
if operation is None:
continue
try:
self._page_transfer(operation)

finally:
self.prefetch_sync_queue.put(
PrefetchAck(
rid=operation.request_id,
completed_req=True,
operation=operation,
)
)
except Empty:
continue

def prefetch_rate_limited(self) -> bool:
"""
Expand All @@ -1194,6 +1215,24 @@ def prefetch_rate_limited(self) -> bool:
# todo: more sophisticated rate limiting based on storage backend performance
return False

def alloc_prefetch_host_buffers(
self, operation: StorageOperation, need_size: int
) -> Optional[torch.Tensor]:
"""Allocate the host bounce for a storage hit."""
return self.mem_pool_host.alloc(need_size)

def can_fit_prefetch_host_buffers(
self, operation: StorageOperation, need_size: int
) -> bool:
"""Whether a prefetch bounce can fit when its host pools are empty."""
return need_size <= self.mem_pool_host.size

def free_prefetch_host_buffers(
self, operation: StorageOperation, host_indices: torch.Tensor
) -> None:
"""Roll back a hit-sized host bounce before transfer ownership moves."""
self.mem_pool_host.free(host_indices)

def _storage_hit_query(self, operation) -> tuple[list[str], int]:
last_hash = operation.last_hash
tokens_to_fetch = operation.token_ids
Expand Down Expand Up @@ -1228,10 +1267,22 @@ def prefetch_thread_func(self):
operation = self.prefetch_queue.get(block=True, timeout=1)
if operation is None:
continue
if operation.is_terminated():
try:
if operation.is_terminated():
hash_value, storage_hit_count = [], 0
else:
hash_value, storage_hit_count = self._storage_hit_query(
operation
)
except Exception:
if not get_memory().enable_unified_memory:
raise
logger.exception(
"HiCache storage query failed for request %s",
operation.request_id,
)
hash_value, storage_hit_count = [], 0
else:
hash_value, storage_hit_count = self._storage_hit_query(operation)

storage_hit_count_tensor = torch.tensor(
storage_hit_count, dtype=torch.int
)
Expand Down
5 changes: 5 additions & 0 deletions python/sglang/srt/managers/schedule_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,7 @@ def match_prefix_for_req(
*,
cow_mamba: bool = False,
include_req: bool = False,
max_prefix_len: Optional[int] = None,
):
if token_ids is None:
token_ids = req.origin_input_ids + req.output_ids
Expand All @@ -155,6 +156,10 @@ def match_prefix_for_req(
# this request's SWA ring. No-op for other layouts.
reprefill_tail = tree_cache.swa_reprefill_tail_tokens()
key_limit = max(0, len(token_ids) - reprefill_tail) if reprefill_tail else None
if max_prefix_len is not None:
key_limit = (
max_prefix_len if key_limit is None else min(key_limit, max_prefix_len)
)

match_result = tree_cache.match_prefix(
MatchPrefixParams(
Expand Down
14 changes: 10 additions & 4 deletions python/sglang/srt/mem_cache/allocator/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from __future__ import annotations

import abc
from typing import TYPE_CHECKING, Protocol
from typing import TYPE_CHECKING, Hashable, Protocol

import torch

Expand Down Expand Up @@ -178,9 +178,11 @@ def merge_and_sort_free(self):
def translate_kv_indices_for_transfer(
self, kv_indices: torch.Tensor
) -> torch.Tensor:
"""Token ids as the PD transfer engine addresses them. Identity here
because a static pool's ids index its registered buffers directly;
virtual-id pools must override."""
"""Token ids as device transfer engines address them.

Identity here: a static pool's token ids index its registered buffers
directly. Virtual-id pools must override.
"""
return kv_indices

def get_cpu_copy(self, indices, mamba_indices=None, req_pool_index=None):
Expand All @@ -191,6 +193,10 @@ def load_cpu_copy(
):
raise NotImplementedError()

def set_hicache_transfer_done_event(self, transfer_key: Hashable, event) -> None:
"""Record an asynchronous HiCache transfer completion event if needed."""
return

def alloc_extend(self, *args, **kwargs):
raise NotImplementedError("alloc_extend is only for paged allocator")

Expand Down
Loading
Loading