[PD]: Support HiCache prefetching and pd-incremental transfer on decode side (#26227)

Co-authored-by: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Zhangheng
2026-06-02 15:40:10 +08:00
committed by GitHub
co-authored by huangtingwei Shangming Cai 晟海
parent 2582134a59
commit 3e993f6140
7 changed files with 630 additions and 56 deletions
+122 -46
View File
@@ -25,7 +25,7 @@ import time
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import numpy as np import numpy as np
import torch import torch
@@ -36,6 +36,13 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.disaggregation.base import KVPoll from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.base.conn import StateType from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver
from sglang.srt.disaggregation.decode_hicache_mixin import (
DecodeHiCachePreallocMixin,
DecodeHiCacheTransferMixin,
DecodePrefixMatch,
HiCacheRestoreGatedKVReceiver,
HiCacheRestoreResult,
)
from sglang.srt.disaggregation.utils import ( from sglang.srt.disaggregation.utils import (
DisaggregationMode, DisaggregationMode,
KVClassType, KVClassType,
@@ -56,7 +63,10 @@ from sglang.srt.managers.schedule_batch import FINISH_ABORT, ScheduleBatch
from sglang.srt.managers.schedule_policy import match_prefix_for_req from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
EvictParams,
)
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
kv_to_page_indices, kv_to_page_indices,
page_align_floor, page_align_floor,
@@ -241,6 +251,13 @@ class DecodeRequest:
waiting_for_input: bool = False waiting_for_input: bool = False
metadata_buffer_index: int = -1 metadata_buffer_index: int = -1
# HiCache Status
prefix_match: Optional[DecodePrefixMatch] = None
hicache_restored_kv_indices: Optional[torch.Tensor] = None
hicache_restored_node: Any = None
hicache_load_consumer_index: int = -1
hicache_restore_status: HiCacheRestoreResult = HiCacheRestoreResult.PENDING
@property @property
def seqlen(self) -> int: def seqlen(self) -> int:
return self.req.seqlen return self.req.seqlen
@@ -250,7 +267,7 @@ class DecodeRequest:
return self.req.priority return self.req.priority
class DecodePreallocQueue: class DecodePreallocQueue(DecodeHiCachePreallocMixin):
""" """
Store the requests that are preallocating. Store the requests that are preallocating.
""" """
@@ -458,7 +475,7 @@ class DecodePreallocQueue:
self.pending_reqs.append(decode_req) self.pending_reqs.append(decode_req)
def _match_prefix_and_lock(self, req: Req) -> Tuple[torch.Tensor, int]: def _match_prefix_and_lock(self, req: Req) -> DecodePrefixMatch:
""" """
Match a request against the decode-side radix cache, lock the matched Match a request against the decode-side radix cache, lock the matched
node to prevent eviction, and return the matched prefix information. node to prevent eviction, and return the matched prefix information.
@@ -470,12 +487,9 @@ class DecodePreallocQueue:
cow_mamba=self.tree_cache.supports_mamba(), cow_mamba=self.tree_cache.supports_mamba(),
include_req=True, include_req=True,
) )
prefix_indices = result.device_indices
last_device_node = result.last_device_node
# Always lock to match aggregated scheduling behavior # Always lock to match aggregated scheduling behavior
self.tree_cache.inc_lock_ref(last_device_node) self.tree_cache.inc_lock_ref(result.last_device_node)
return self._build_decode_prefix_match(req, result)
return prefix_indices, len(prefix_indices)
def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]: def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]:
prefill_info = self.kv_manager.prefill_info_table.get(_bootstrap_addr(req)) prefill_info = self.kv_manager.prefill_info_table.get(_bootstrap_addr(req))
@@ -756,7 +770,8 @@ class DecodePreallocQueue:
full_allocatable_tokens = self._allocatable_token_budgets( full_allocatable_tokens = self._allocatable_token_budgets(
retractable_tokens=retractable_tokens, count_retracted=True retractable_tokens=retractable_tokens, count_retracted=True
) )
reserved_restore_tokens = self._hicache_pending_restore_tokens()
full_allocatable_tokens -= reserved_restore_tokens
# Sort by priority before any index-based bookkeeping so that both the # Sort by priority before any index-based bookkeeping so that both the
# abort-scan loop and the preallocation loop operate on the same order. # abort-scan loop and the preallocation loop operate on the same order.
if self.scheduler.enable_priority_scheduling: if self.scheduler.enable_priority_scheduling:
@@ -815,15 +830,18 @@ class DecodePreallocQueue:
# Memory estimation: don't add if the projected memory cannot be met # Memory estimation: don't add if the projected memory cannot be met
# TODO: add new_token ratio # TODO: add new_token ratio
origin_input_len = len(decode_req.req.origin_input_ids) origin_input_len = len(decode_req.req.origin_input_ids)
prefix_match: Optional[DecodePrefixMatch] = None
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache: if self.scheduler.server_args.disaggregation_decode_enable_radix_cache:
# Match prefix against decode's radix cache. # Match prefix against decode's radix cache.
prefix_indices, prefix_len = self._match_prefix_and_lock(decode_req.req) prefix_match = self._match_prefix_and_lock(decode_req.req)
# Align prefix_len down to page boundary so both prefill and prefix_indices = prefix_match.prefix_indices
# decode agree on the page-aligned split point for KV transfer. # prefix_len: tokens already on device (L1 hit).
page_size = self.token_to_kv_pool_allocator.page_size # total_prefix_len: full prefix promised to prefill
if page_size > 1 and prefix_len % page_size != 0: # (L1 + L2 host hit + L3 storage hit), sent as PD
prefix_len = page_align_floor(prefix_len, page_size) # protocol's `decode_prefix_len`. The [prefix_len, total)
prefix_indices = prefix_indices[:prefix_len] # gap is filled by HiCache loadback later.
prefix_len = prefix_match.l1_prefix_len
total_prefix_len = prefix_match.decode_prefix_len
fill_len = origin_input_len + max(len(decode_req.req.output_ids) - 1, 0) fill_len = origin_input_len + max(len(decode_req.req.output_ids) - 1, 0)
required_alloc_tokens = self._required_alloc_tokens( required_alloc_tokens = self._required_alloc_tokens(
@@ -836,10 +854,12 @@ class DecodePreallocQueue:
retractable_tokens=retractable_tokens, retractable_tokens=retractable_tokens,
count_retracted=True, count_retracted=True,
extra_reserved_reqs=len(preallocated_reqs), extra_reserved_reqs=len(preallocated_reqs),
hicache_reserved_tokens=reserved_restore_tokens,
) )
else: else:
prefix_indices = None prefix_indices = None
prefix_len = 0 prefix_len = 0
total_prefix_len = 0
required_alloc_tokens = origin_input_len required_alloc_tokens = origin_input_len
required_tokens_for_request = ( required_tokens_for_request = (
@@ -885,20 +905,31 @@ class DecodePreallocQueue:
self.tree_cache.dec_lock_ref(decode_req.req.last_node) self.tree_cache.dec_lock_ref(decode_req.req.last_node)
break break
dst_kv_indices = self._pre_alloc(decode_req.req, prefix_indices, prefix_len) dst_kv_indices = self._pre_alloc(
decode_req.req,
prefix_indices,
prefix_len,
total_prefix_len,
)
decode_req.prefix_match = prefix_match
if self.scheduler.enable_decode_hicache:
self._start_hicache_prefetch(decode_req.req, prefix_match)
hisparse_req_budget -= 1 hisparse_req_budget -= 1
# Recompute from actual pool state for the next queue entry. # Recompute from actual pool state for the next queue entry.
# This accounts for page rounding and newly locked evictable cache. # This accounts for page rounding and newly locked evictable cache.
if prefix_match is not None:
reserved_restore_tokens += prefix_match.restore_token_count
full_allocatable_tokens = self._allocatable_token_budgets( full_allocatable_tokens = self._allocatable_token_budgets(
retractable_tokens=retractable_tokens, retractable_tokens=retractable_tokens,
count_retracted=True, count_retracted=True,
extra_reserved_reqs=len(preallocated_reqs) + 1, extra_reserved_reqs=len(preallocated_reqs) + 1,
hicache_reserved_tokens=reserved_restore_tokens,
) )
if uses_swa_tail_prealloc: if uses_swa_tail_prealloc:
# SWA budget uses simple decrement (no radix cache eviction in # SWA budget uses simple decrement (no radix cache eviction in
# the SWA pool, so page-rounding drift is negligible). # the SWA pool, so page-rounding drift is negligible).
swa_allocatable_tokens -= swa_required swa_allocatable_tokens -= swa_required
decode_req.req.cache_protected_len = prefix_len decode_req.req.cache_protected_len = total_prefix_len
page_size = self.token_to_kv_pool_allocator.page_size page_size = self.token_to_kv_pool_allocator.page_size
if self.scheduler.enable_hisparse: if self.scheduler.enable_hisparse:
@@ -913,7 +944,7 @@ class DecodePreallocQueue:
# Only send delta indices (beyond prefix) to prefill. # Only send delta indices (beyond prefix) to prefill.
kv_indices = ( kv_indices = (
self.req_to_token_pool.req_to_token[decode_req.req.req_pool_idx][ self.req_to_token_pool.req_to_token[decode_req.req.req_pool_idx][
prefix_len:origin_input_len total_prefix_len:origin_input_len
] ]
.cpu() .cpu()
.numpy() .numpy()
@@ -977,7 +1008,7 @@ class DecodePreallocQueue:
page_indices, page_indices,
decode_req.metadata_buffer_index, decode_req.metadata_buffer_index,
state_indices, state_indices,
decode_prefix_len=prefix_len, decode_prefix_len=total_prefix_len,
) )
if ( if (
self.transfer_queue.enable_staging self.transfer_queue.enable_staging
@@ -1065,6 +1096,7 @@ class DecodePreallocQueue:
count_retracted: bool = True, count_retracted: bool = True,
extra_reserved_reqs: int = 0, extra_reserved_reqs: int = 0,
reserved_tokens: Optional[int] = None, reserved_tokens: Optional[int] = None,
hicache_reserved_tokens: int = 0,
) -> int: ) -> int:
need_space_for_single_req = self._need_space_for_single_req(retractable_tokens) need_space_for_single_req = self._need_space_for_single_req(retractable_tokens)
if reserved_tokens is None: if reserved_tokens is None:
@@ -1107,6 +1139,7 @@ class DecodePreallocQueue:
full_required, _ = self._prealloc_required_tokens(req) full_required, _ = self._prealloc_required_tokens(req)
allocatable_tokens -= full_required allocatable_tokens -= full_required
allocatable_tokens -= hicache_reserved_tokens
return allocatable_tokens return allocatable_tokens
def _swa_tail_allocatable_token_budget( def _swa_tail_allocatable_token_budget(
@@ -1189,10 +1222,20 @@ class DecodePreallocQueue:
req: Req, req: Req,
prefix_indices: Optional[torch.Tensor] = None, prefix_indices: Optional[torch.Tensor] = None,
prefix_len: Optional[int] = None, prefix_len: Optional[int] = None,
total_prefix_len: Optional[int] = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""Pre-allocate the memory for req_to_token and token_kv_pool""" """Pre-allocate the memory for req_to_token and token_kv_pool.
``prefix_len`` is the L1 device-resident prefix length (already
backed by ``prefix_indices``). ``total_prefix_len`` is the full
prefix committed to prefill as ``decode_prefix_len`` (L1 + L2 + L3);
the ``[prefix_len, total_prefix_len)`` gap is filled later by HiCache
loadback.
"""
if prefix_len is None: if prefix_len is None:
prefix_len = 0 prefix_len = 0
if total_prefix_len is None:
total_prefix_len = prefix_len
req_pool_indices = self.req_to_token_pool.alloc([req]) req_pool_indices = self.req_to_token_pool.alloc([req])
@@ -1212,7 +1255,7 @@ class DecodePreallocQueue:
# TODO(retraction): when retraction is implemented with radix cache # TODO(retraction): when retraction is implemented with radix cache
# awareness, a retracted request should re-match the tree here # awareness, a retracted request should re-match the tree here
# instead of re-allocating from scratch. See resume_retracted_reqs. # instead of re-allocating from scratch. See resume_retracted_reqs.
delta_len = fill_len - prefix_len delta_len = fill_len - total_prefix_len
required_alloc_tokens = self._required_alloc_tokens( required_alloc_tokens = self._required_alloc_tokens(
fill_len=fill_len, prefix_len=prefix_len fill_len=fill_len, prefix_len=prefix_len
) )
@@ -1233,7 +1276,8 @@ class DecodePreallocQueue:
f"after evicting {result.num_tokens_evicted}/{num_to_evict} tokens. " f"after evicting {result.num_tokens_evicted}/{num_to_evict} tokens. "
f"evictable_size={self.tree_cache.evictable_size()}, " f"evictable_size={self.tree_cache.evictable_size()}, "
f"protected_size={self.tree_cache.protected_size()}, " f"protected_size={self.tree_cache.protected_size()}, "
f"fill_len={fill_len}, prefix_len={prefix_len}, delta_len={delta_len}, " f"fill_len={fill_len}, prefix_len={prefix_len}, "
f"total_prefix_len={total_prefix_len}, delta_len={delta_len}, "
f"page_size={self.token_to_kv_pool_allocator.page_size}, " f"page_size={self.token_to_kv_pool_allocator.page_size}, "
f"req={req.rid}" f"req={req.rid}"
) )
@@ -1289,9 +1333,9 @@ class DecodePreallocQueue:
else: else:
kv_loc = self.token_to_kv_pool_allocator.alloc_extend( kv_loc = self.token_to_kv_pool_allocator.alloc_extend(
prefix_lens=torch.tensor( 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=torch.tensor([fill_len], dtype=torch.int64, device=device),
seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64), seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64),
last_loc=last_loc, last_loc=last_loc,
@@ -1304,13 +1348,17 @@ class DecodePreallocQueue:
f"evictable={self.tree_cache.evictable_size()}, " f"evictable={self.tree_cache.evictable_size()}, "
f"protected={self.tree_cache.protected_size()}, " f"protected={self.tree_cache.protected_size()}, "
f"required_alloc={required_alloc_tokens}, delta={delta_len}, " f"required_alloc={required_alloc_tokens}, delta={delta_len}, "
f"fill={fill_len}, prefix={prefix_len}, " f"fill={fill_len}, prefix={prefix_len}, total_prefix={total_prefix_len}, "
f"page_size={self.token_to_kv_pool_allocator.page_size}, " f"page_size={self.token_to_kv_pool_allocator.page_size}, "
f"req={req.rid}" f"req={req.rid}"
) )
self.req_to_token_pool.write( self.req_to_token_pool.write(
(req.req_pool_idx, slice(prefix_len, prefix_len + len(kv_loc))), kv_loc (
req.req_pool_idx,
slice(total_prefix_len, total_prefix_len + len(kv_loc)),
),
kv_loc,
) )
# Truncate fill_ids to kv_committed_len so cache_unfinished_req only # Truncate fill_ids to kv_committed_len so cache_unfinished_req only
@@ -1324,7 +1372,7 @@ class DecodePreallocQueue:
req.prefix_indices = ( req.prefix_indices = (
prefix_indices if prefix_len > 0 else torch.empty((0,), dtype=torch.int64) prefix_indices if prefix_len > 0 else torch.empty((0,), dtype=torch.int64)
) )
req.set_extend_input_len(len(req.fill_ids) - prefix_len) req.set_extend_input_len(len(req.fill_ids) - total_prefix_len)
# Return the transfer destination indices: # Return the transfer destination indices:
if self.scheduler.enable_hisparse: if self.scheduler.enable_hisparse:
@@ -1332,7 +1380,7 @@ class DecodePreallocQueue:
return kv_loc return kv_loc
class DecodeTransferQueue: class DecodeTransferQueue(DecodeHiCacheTransferMixin):
""" """
Store the requests that is polling kv Store the requests that is polling kv
""" """
@@ -1433,7 +1481,9 @@ class DecodeTransferQueue:
decode_req.kv_receiver = None decode_req.kv_receiver = None
return return
# Success - commit the transfer self._commit_hicache_local_restore_to_req(decode_req)
# Case 3: Success - commit the transfer
decode_req.req.output_ids.append(output_id[0].item()) decode_req.req.output_ids.append(output_id[0].item())
decode_req.req.cached_tokens = cached_tokens[0].item() decode_req.req.cached_tokens = cached_tokens[0].item()
# The prefill node already reported its prefix-cache hit in # The prefill node already reported its prefix-cache hit in
@@ -1475,8 +1525,13 @@ class DecodeTransferQueue:
return return
def _poll_with_metadata_gate(self) -> List[int]: def _poll_with_metadata_gate(self) -> List[int]:
pollers = (
[HiCacheRestoreGatedKVReceiver(dr) for dr in self.queue]
if self.scheduler.enable_decode_hicache
else [dr.kv_receiver for dr in self.queue]
)
return poll_and_all_reduce( return poll_and_all_reduce(
[dr.kv_receiver for dr in self.queue], pollers,
self.gloo_group, self.gloo_group,
decode_reqs=self.queue, decode_reqs=self.queue,
metadata_buffers=self.metadata_buffers, metadata_buffers=self.metadata_buffers,
@@ -1507,6 +1562,15 @@ class DecodeTransferQueue:
if not self.queue: if not self.queue:
return [] return []
if self.scheduler.enable_decode_hicache:
self._process_hicache_local_restores(
[
decode_req
for decode_req in self.queue
if rids_to_check is None or decode_req.req.rid in rids_to_check
]
)
if self.enable_staging: if self.enable_staging:
polls = self._poll_with_staging() polls = self._poll_with_staging()
else: else:
@@ -1518,12 +1582,21 @@ class DecodeTransferQueue:
if rids_to_check is not None and decode_req.req.rid not in rids_to_check: if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
continue continue
hicache_restore_status = decode_req.hicache_restore_status
if (
poll == KVPoll.Failed
or hicache_restore_status == HiCacheRestoreResult.FAILED
):
error_message = (
f"Decode transfer failed for request rank={self.tp_rank} "
f"{decode_req.req.rid=} {decode_req.req.bootstrap_room=}"
)
if poll == KVPoll.Failed: if poll == KVPoll.Failed:
error_message = f"Decode transfer failed for request rank={self.tp_rank} {decode_req.req.rid=} {decode_req.req.bootstrap_room=}"
try: try:
decode_req.kv_receiver.failure_exception() decode_req.kv_receiver.failure_exception()
except Exception as e: except Exception as e:
error_message += f" with exception {e}" error_message += f" with exception {e}"
self._clean_hicache_prefetch_resources(decode_req)
logger.error(error_message) logger.error(error_message)
prepare_abort( prepare_abort(
decode_req.req, decode_req.req,
@@ -1543,6 +1616,11 @@ class DecodeTransferQueue:
self.scheduler.metrics_collector.increment_transfer_failed_reqs() self.scheduler.metrics_collector.increment_transfer_failed_reqs()
continue continue
elif poll == KVPoll.Success: elif poll == KVPoll.Success:
if (
self.scheduler.enable_decode_hicache
and hicache_restore_status == HiCacheRestoreResult.PENDING
):
continue
self._commit_transfer_to_req(decode_req) self._commit_transfer_to_req(decode_req)
indices_to_remove.add(i) indices_to_remove.add(i)
# Check if request was aborted due to corruption # Check if request was aborted due to corruption
@@ -1555,6 +1633,7 @@ class DecodeTransferQueue:
self.scheduler.hisparse_coordinator.request_finished( self.scheduler.hisparse_coordinator.request_finished(
decode_req.req decode_req.req
) )
self._clean_hicache_prefetch_resources(decode_req)
release_kv_cache(decode_req.req, self.tree_cache, is_insert=False) release_kv_cache(decode_req.req, self.tree_cache, is_insert=False)
if self.scheduler.metrics_reporter.enable_metrics: if self.scheduler.metrics_reporter.enable_metrics:
self.scheduler.metrics_collector.increment_transfer_failed_reqs() self.scheduler.metrics_collector.increment_transfer_failed_reqs()
@@ -1730,19 +1809,13 @@ class SchedulerDisaggregationDecodeMixin:
# we can only add at least `num_not_used_batch` new batch to the running queue # we can only add at least `num_not_used_batch` new batch to the running queue
if i < num_not_used_batch: if i < num_not_used_batch:
can_run_list.append(req) can_run_list.append(req)
# Decode-radix path: do NOT re-match prefix here. # Decode-radix path: new requests already matched in
# `pop_preallocated` already took a tree snapshot and used it # `pop_preallocated`. Retracted requests reset `last_node`,
# to (1) pre-allocate KV, (2) choose delta pages for transfer, # so re-match only when that state is missing.
# and (3) set cache_protected_len/last_node for correct frees. if self.server_args.disaggregation_decode_enable_radix_cache:
# Re-matching now can observe a newer tree (other reqs may have tree_cache = self.tree_cache if req.last_node is None else None
# inserted the same prefix) and overwrite cache_protected_len, else:
# making `cache_unfinished_req` free the wrong range (leak). tree_cache = self.tree_cache
# Non-radix decode keeps the original behavior.
tree_cache = (
None
if self.server_args.disaggregation_decode_enable_radix_cache
else self.tree_cache
)
req.init_next_round_input(tree_cache) req.init_next_round_input(tree_cache)
# Truncate fill_ids to kv_committed_len so cache_unfinished_req # Truncate fill_ids to kv_committed_len so cache_unfinished_req
# only sees committed KV (fill_ids includes one uncommitted token). # only sees committed KV (fill_ids includes one uncommitted token).
@@ -1778,6 +1851,9 @@ class SchedulerDisaggregationDecodeMixin:
return new_batch return new_batch
def process_decode_queue(self: Scheduler): def process_decode_queue(self: Scheduler):
if self.enable_decode_hicache:
self.tree_cache.check_hicache_events()
if self.server_args.disaggregation_decode_enable_offload_kvcache: if self.server_args.disaggregation_decode_enable_offload_kvcache:
self.decode_offload_manager.check_offload_progress() self.decode_offload_manager.check_offload_progress()
@@ -0,0 +1,311 @@
"""HiCache integration mixins for the decode side of PD disaggregation"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from enum import Enum
from typing import TYPE_CHECKING, Any, List, Optional
import torch
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.managers.schedule_policy import match_prefix_for_req
from sglang.srt.mem_cache.base_prefix_cache import InitLoadBackParams
if TYPE_CHECKING:
from sglang.srt.disaggregation.decode import DecodeRequest
from sglang.srt.managers.schedule_batch import Req
logger = logging.getLogger(__name__)
@dataclass
class DecodePrefixMatch:
prefix_indices: torch.Tensor
l2_host_hit_length: int
l3_storage_hit_length: int
last_device_node: Any
last_host_node: Any = None
prefetch_registered: bool = False
@property
def l1_prefix_len(self) -> int:
return len(self.prefix_indices)
@property
def decode_prefix_len(self) -> int:
return self.l1_prefix_len + self.l2_host_hit_length + self.l3_storage_hit_length
@property
def needs_local_restore(self) -> bool:
return self.decode_prefix_len > self.l1_prefix_len
@property
def restore_token_count(self) -> int:
"""Number of tokens that need L2/L3 load_back to device."""
return self.decode_prefix_len - self.l1_prefix_len
class HiCacheRestoreResult(Enum):
"""Outcome of one tick of the HiCache local-restore state machine."""
PENDING = "pending"
READY = "ready"
FAILED = "failed"
class DecodeHiCachePreallocMixin:
"""HiCache hooks for ``DecodePreallocQueue``: issue prefetch + reserve tokens."""
def _build_decode_prefix_match(self, req: "Req", result: Any) -> DecodePrefixMatch:
"""Convert a ``match_prefix_for_req`` result into ``DecodePrefixMatch``.
Performs the optional L3 storage hit length query when decode-side
HiCache is enabled and the last host node is backed up.
"""
prefix_indices = result.device_indices
l1_prefix_len = len(prefix_indices)
l2_host_hit_length = result.host_hit_length
l3_storage_hit_length = 0
last_host_node = None
if self.scheduler.enable_decode_hicache:
last_host_node = result.last_host_node
if last_host_node.backuped or last_host_node is self.tree_cache.root_node:
matched_len = l1_prefix_len + l2_host_hit_length
suffix_tokens = req.origin_input_ids[matched_len:]
last_hash = last_host_node.get_last_hash_value()
prefix_keys = (
last_host_node.get_prefix_hash_values(last_host_node.parent)
if self.tree_cache.hicache_storage_pass_prefix_keys
else None
)
l3_storage_hit_length = self.tree_cache.query_storage_hit_length(
last_host_node,
suffix_tokens,
last_hash,
prefix_keys,
)
return DecodePrefixMatch(
prefix_indices=prefix_indices,
l2_host_hit_length=l2_host_hit_length,
l3_storage_hit_length=l3_storage_hit_length,
last_device_node=result.last_device_node,
last_host_node=last_host_node if l3_storage_hit_length > 0 else None,
)
def _start_hicache_prefetch(
self, req: "Req", prefix_match: Optional["DecodePrefixMatch"]
) -> None:
"""Issue L3 storage prefetch after admission succeeds.
On failure, degrades to L2-only restore by clearing l3 fields.
"""
if (
prefix_match is None
or prefix_match.l3_storage_hit_length <= 0
or prefix_match.last_host_node is None
):
return
try:
node = prefix_match.last_host_node
matched_len = prefix_match.l1_prefix_len + prefix_match.l2_host_hit_length
suffix = req.origin_input_ids[
matched_len : matched_len + prefix_match.l3_storage_hit_length
]
last_hash = node.get_last_hash_value()
prefix_keys = (
node.get_prefix_hash_values(node.parent)
if self.tree_cache.hicache_storage_pass_prefix_keys
else None
)
self.tree_cache.prefetch_from_storage(
req.rid, node, suffix, last_hash, prefix_keys
)
prefix_match.prefetch_registered = (
req.rid in self.tree_cache.ongoing_prefetch
)
except Exception as e:
logger.warning(
"HiCache L3 prefetch failed for rid=%s: %s; falling back to L2-only LoadingBack",
req.rid,
e,
)
prefix_match.l3_storage_hit_length = 0
prefix_match.prefetch_registered = False
def _hicache_pending_restore_tokens(self) -> int:
"""Total device tokens reserved for pending HiCache L2/L3 load_back."""
if not self.scheduler.enable_decode_hicache:
return 0
return sum(
dr.prefix_match.restore_token_count
for dr in self.transfer_queue.queue
if dr.prefix_match is not None
and dr.hicache_restore_status == HiCacheRestoreResult.PENDING
and dr.hicache_restored_node is None
)
class HiCacheRestoreGatedKVReceiver:
"""Wraps a kv_receiver so KVPoll.Success is gated on HiCache restore READY."""
def __init__(self, decode_req: "DecodeRequest"):
self.decode_req = decode_req
def poll(self) -> KVPoll:
poll = self.decode_req.kv_receiver.poll()
if (
poll == KVPoll.Success
and self.decode_req.hicache_restore_status == HiCacheRestoreResult.PENDING
):
return KVPoll.Transferring
return poll
class DecodeHiCacheTransferMixin:
"""HiCache hooks for ``DecodeTransferQueue``: drive restore state machine."""
def _clean_hicache_prefetch_resources(self, decode_req: "DecodeRequest") -> None:
if (
decode_req.prefix_match is not None
and decode_req.prefix_match.prefetch_registered
):
self.tree_cache.release_aborted_request(decode_req.req.rid)
if decode_req.hicache_restored_node is not None:
self.tree_cache.dec_lock_ref(decode_req.hicache_restored_node)
decode_req.hicache_restored_node = None
def _try_hicache_queue_load_back(self, dr: "DecodeRequest") -> bool:
"""Queue one L2->L1 load_back op for ``dr``; True iff a DMA was queued.
On success, ``dr.hicache_restored_node`` and ``hicache_restored_kv_indices``
are populated, and an inc_lock_ref is held until commit/abort.
Trivial cases (all-on-device / no needed coverage) auto-flip to READY.
Failback paths flip to FAILED.
"""
pm = dr.prefix_match
# Wait for L3 -> L2 prefetch to drain (skip when no L3 hit).
if pm.l3_storage_hit_length > 0:
if not self.tree_cache.check_prefetch_progress(dr.req.rid):
return False
self.tree_cache.pop_prefetch_loaded_tokens(dr.req.rid)
# Re-match: req.last_node / prefix_indices updated to current device state.
rematch = match_prefix_for_req(
self.tree_cache,
dr.req,
dr.req.origin_input_ids,
cow_mamba=False,
include_req=True,
)
new_indices, restored_node = self.tree_cache.init_load_back(
InitLoadBackParams(
best_match_node=rematch.best_match_node,
host_hit_length=rematch.host_hit_length,
req=dr.req,
)
)
# Failback: total coverage < required prefix means device alloc likely failed.
if len(rematch.device_indices) + len(new_indices) < pm.decode_prefix_len:
logger.warning(
"HiCache load_back failed for rid=%s: device_indices=%d, "
"new_indices=%d, expected decode_prefix_len=%d (l1=%d, l2=%d, l3=%d)",
dr.req.rid,
len(rematch.device_indices),
len(new_indices),
pm.decode_prefix_len,
pm.l1_prefix_len,
pm.l2_host_hit_length,
pm.l3_storage_hit_length,
)
dr.hicache_restore_status = HiCacheRestoreResult.FAILED
return False
dr.hicache_restored_kv_indices = torch.cat(
[rematch.device_indices[pm.l1_prefix_len :], new_indices]
)
dr.hicache_restored_node = restored_node
self.tree_cache.inc_lock_ref(restored_node)
if len(new_indices) == 0:
# Whole prefix already on device; no DMA needed.
dr.hicache_restore_status = HiCacheRestoreResult.READY
return False
return True
def _process_hicache_local_restores(
self, decode_reqs: List["DecodeRequest"]
) -> None:
if not hasattr(self.tree_cache, "is_load_back_event_done"):
return
# Filter once: keep only PENDING reqs that still need restore work;
# trivially-done reqs (no prefix_match / nothing to restore) flip to READY.
active: List["DecodeRequest"] = []
for dr in decode_reqs:
if dr.hicache_restore_status != HiCacheRestoreResult.PENDING:
continue
pm = dr.prefix_match
if pm is None or not pm.needs_local_restore:
dr.hicache_restore_status = HiCacheRestoreResult.READY
continue
active.append(dr)
# Phase A: advance in-flight DMAs to READY.
for dr in active:
if (
dr.hicache_restored_node is not None
and self.tree_cache.is_load_back_event_done(
dr.hicache_load_consumer_index
)
):
dr.hicache_restore_status = HiCacheRestoreResult.READY
# Phase B: queue new load_back ops if the next slot is free.
# The (producer_index + 1) check ensures we never overwrite a still-in-flight slot:
# if a previous req holds that slot and isn't done, its event won't be signaled.
counter = self.tree_cache.cache_controller.layer_done_counter
if not self.tree_cache.is_load_back_event_done(
(counter.producer_index + 1) % counter.num_counters
):
return
queued = [
dr
for dr in active
if dr.hicache_restored_node is None
and self._try_hicache_queue_load_back(dr)
]
if not queued:
return
# Phase C: kick off merged DMA, bind consumer_index for Phase A polling next tick.
consumer_index = self.tree_cache.ready_to_load_host_cache()
if consumer_index < 0:
for dr in queued:
dr.hicache_restore_status = HiCacheRestoreResult.READY
return
for dr in queued:
dr.hicache_load_consumer_index = consumer_index
def _commit_hicache_local_restore_to_req(self, decode_req: "DecodeRequest") -> None:
prefix_match = decode_req.prefix_match
if prefix_match is None or not prefix_match.needs_local_restore:
return
self.tree_cache.dec_lock_ref(prefix_match.last_device_node)
self.tree_cache.req_to_token_pool.write(
(
decode_req.req.req_pool_idx,
slice(prefix_match.l1_prefix_len, prefix_match.decode_prefix_len),
),
decode_req.hicache_restored_kv_indices,
)
decode_req.req.prefix_indices = torch.cat(
[prefix_match.prefix_indices, decode_req.hicache_restored_kv_indices]
)
decode_req.req.last_node = decode_req.hicache_restored_node
+4
View File
@@ -339,6 +339,10 @@ class Scheduler(
self.page_size = server_args.page_size self.page_size = server_args.page_size
self.enable_hierarchical_cache = server_args.enable_hierarchical_cache self.enable_hierarchical_cache = server_args.enable_hierarchical_cache
self.enable_hicache_storage = server_args.hicache_storage_backend is not None self.enable_hicache_storage = server_args.hicache_storage_backend is not None
self.enable_decode_hicache = (
server_args.disaggregation_decode_enable_radix_cache
and self.enable_hierarchical_cache
)
self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get() self.max_recv_per_poll = envs.SGLANG_SCHEDULER_MAX_RECV_PER_POLL.get()
self.enable_hisparse = server_args.enable_hisparse self.enable_hisparse = server_args.enable_hisparse
self.hisparse_coordinator: Optional[HiSparseCoordinator] = None self.hisparse_coordinator: Optional[HiSparseCoordinator] = None
@@ -852,6 +852,20 @@ class HiRadixCache(RadixCache):
# ACK until all events are processed # ACK until all events are processed
del self.cache_controller.ack_load_queue[:finish_count] del self.cache_controller.ack_load_queue[:finish_count]
def is_load_back_event_done(self, consumer_index: int) -> bool:
"""Return True after the local load-back event is complete."""
if consumer_index < 0:
return True
finish_event = self.cache_controller.layer_done_counter.events[
consumer_index
].finish_event
if not finish_event.query():
return False
self.loading_check()
return True
def evictable_size(self): def evictable_size(self):
return self.evictable_size_ return self.evictable_size_
@@ -1110,6 +1124,42 @@ class HiRadixCache(RadixCache):
last_node, last_node,
) )
def query_storage_hit_length(
self,
last_host_node: TreeNode,
new_input_tokens: List[int],
last_hash: Optional[str] = None,
prefix_keys: Optional[List[str]] = None,
) -> int:
if not self.enable_storage or self.cache_controller.prefetch_rate_limited():
return 0
prefetch_key = RadixKey(
new_input_tokens,
extra_key=last_host_node.key.extra_key,
is_bigram=self.is_eagle,
).page_aligned(self.page_size)
if len(prefetch_key) < self.prefetch_threshold:
return 0
operation = PrefetchOperation(
"__storage_hit_query__",
self.cache_controller.mem_pool_host.get_dummy_flat_data_page()[:0],
prefetch_key,
last_hash,
prefix_keys,
)
hash_values, storage_hit_count = self.cache_controller._storage_hit_query(
operation
)
storage_hit_count_tensor = torch.tensor(storage_hit_count, dtype=torch.int)
self._all_reduce_attn_groups(
storage_hit_count_tensor, torch.distributed.ReduceOp.MIN
)
storage_hit_count = storage_hit_count_tensor.item()
storage_hit_count = storage_hit_count - (storage_hit_count % self.page_size)
return storage_hit_count
def ready_to_load_host_cache(self) -> int: def ready_to_load_host_cache(self) -> int:
""" """
Notify the cache controller to start the KV cache loading. Notify the cache controller to start the KV cache loading.
@@ -1,11 +1,21 @@
import asyncio
import os
import shutil
import tempfile
import time import time
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
import requests import requests
from sglang.benchmark.datasets.random import sample_random_requests
from sglang.benchmark.utils import get_tokenizer
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.cache_hit_kit import run_multiturn_cache_hit_test from sglang.test.kits.cache_hit_kit import (
async_request_sglang_generate,
gen_payload,
run_multiturn_cache_hit_test,
)
from sglang.test.run_eval import run_eval from sglang.test.run_eval import run_eval
from sglang.test.server_fixtures.disaggregation_fixture import ( from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase, PDDisaggregationServerBase,
@@ -140,5 +150,120 @@ class TestDisaggregationDecodeRadixCacheMooncake(
transfer_backend_name = "mooncake" transfer_backend_name = "mooncake"
@unittest.skipUnless(
is_in_ci() or _has_mooncake(),
"Mooncake is required for decode radix cache disaggregation coverage.",
)
class TestDisaggregationDecodeRadixHiCacheFileBackend(PDDisaggregationServerBase):
extra_prefill_args = [
"--enable-hierarchical-cache",
"--hicache-ratio",
"1.2",
"--hicache-write-policy",
"write_through",
"--hicache-storage-backend",
"file",
"--hicache-storage-prefetch-policy",
"wait_complete",
"--hicache-io-backend",
"kernel",
"--hicache-mem-layout",
"page_first",
"--page-size",
"64",
]
extra_decode_args = [
"--disaggregation-decode-enable-radix-cache",
*extra_prefill_args,
]
transfer_backend_name = "mooncake"
@classmethod
def setUpClass(cls):
cls.hicache_dir = tempfile.mkdtemp(prefix="sglang-hicache-")
os.environ["SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR"] = cls.hicache_dir
super().setUpClass()
cls.model = try_cached_model(DEFAULT_MODEL_NAME_FOR_TEST)
cls.transfer_backend = [
"--disaggregation-transfer-backend",
cls.transfer_backend_name,
]
cls.launch_all()
@classmethod
def tearDownClass(cls):
super().tearDownClass()
os.environ.pop("SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR", None)
shutil.rmtree(cls.hicache_dir, ignore_errors=True)
def _post_ok(self, url):
response = requests.post(url, timeout=60)
response.raise_for_status()
def _flush_memory_cache(self):
self._post_ok(f"{self.prefill_url}/flush_cache?timeout=30")
self._post_ok(f"{self.decode_url}/flush_cache?timeout=30")
def _generate(self, input_ids, output_len):
output = asyncio.run(
async_request_sglang_generate(
gen_payload(input_ids, output_len),
f"{self.base_url}/generate",
)
)
self.assertTrue(output.success, output.error)
return output
def _sample_token_ids(self, input_len, output_len, num_prompts=1):
tokenizer = get_tokenizer(self.model)
return [
list(request.prompt)
for request in sample_random_requests(
input_len=input_len,
output_len=output_len,
num_prompts=num_prompts,
range_ratio=1.0,
tokenizer=tokenizer,
dataset_path="",
return_text=False,
)
]
def test_decode_hicache_file_backend_l3_reuses_decode_output_after_flush(self):
self._post_ok(f"{self.decode_url}/hicache/storage-backend/clear")
self._flush_memory_cache()
num_rounds = 5
output_len = 64
history = self._sample_token_ids(
input_len=256, output_len=output_len, num_prompts=1
)[0]
suffixes = self._sample_token_ids(
input_len=64, output_len=output_len, num_prompts=num_rounds - 1
)
prev_prompt_len = 0
prev_output_len = 0
for round_idx in range(num_rounds):
output = self._generate(history, output_len)
if round_idx == 0:
self.assertEqual(output.cached_tokens, 0)
else:
self.assertGreaterEqual(
output.cached_tokens,
prev_prompt_len + prev_output_len,
)
history.extend(output.output_ids)
prev_prompt_len = output.prompt_len
prev_output_len = len(output.output_ids)
if round_idx < num_rounds - 1:
history.extend(suffixes[round_idx])
time.sleep(1)
self._flush_memory_cache()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -106,11 +106,13 @@ class TestDecodePreallocQueuePriority(unittest.TestCase):
queue._resolve_pending_reqs = MagicMock() queue._resolve_pending_reqs = MagicMock()
queue._update_handshake_waiters = MagicMock() queue._update_handshake_waiters = MagicMock()
queue._allocatable_tokens = MagicMock(return_value=1000) queue._allocatable_tokens = MagicMock(return_value=1000)
queue._pre_alloc = MagicMock(
side_effect=lambda req, prefix_indices=None, prefix_len=0: torch.arange( def pre_alloc_mock(req, prefix_indices=None, prefix_len=0, total_prefix_len=0):
return torch.arange(
len(req.origin_input_ids) - prefix_len, dtype=torch.int64 len(req.origin_input_ids) - prefix_len, dtype=torch.int64
) )
)
queue._pre_alloc = MagicMock(side_effect=pre_alloc_mock)
queue.req_to_token_pool = MagicMock() queue.req_to_token_pool = MagicMock()
queue.req_to_token_pool.available_size.return_value = 100 queue.req_to_token_pool.available_size.return_value = 100
@@ -32,6 +32,7 @@ from unittest.mock import MagicMock
import torch import torch
from sglang.srt.disaggregation.decode import DecodePreallocQueue from sglang.srt.disaggregation.decode import DecodePreallocQueue
from sglang.srt.disaggregation.decode_hicache_mixin import DecodePrefixMatch
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
InsertParams, InsertParams,
MatchPrefixParams, MatchPrefixParams,
@@ -312,7 +313,12 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
queue._resolve_pending_reqs = MagicMock() queue._resolve_pending_reqs = MagicMock()
queue._update_handshake_waiters = MagicMock() queue._update_handshake_waiters = MagicMock()
queue._match_prefix_and_lock = MagicMock( queue._match_prefix_and_lock = MagicMock(
return_value=(torch.arange(4, dtype=torch.int64), 4) return_value=DecodePrefixMatch(
prefix_indices=torch.arange(4, dtype=torch.int64),
l2_host_hit_length=0,
l3_storage_hit_length=0,
last_device_node=req.last_node,
)
) )
queue._pre_alloc = MagicMock( queue._pre_alloc = MagicMock(
side_effect=AssertionError("_pre_alloc should not run") side_effect=AssertionError("_pre_alloc should not run")