[PD] Preserve decode KV across retraction in HiCache (#34801)

Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
cctry
2026-08-17 08:49:11 -07:00
committed by GitHub
co-authored by cctry
parent af743371cc
commit 2e7c85da68
15 changed files with 779 additions and 34 deletions
+15 -2
View File
@@ -75,6 +75,8 @@ from sglang.srt.mem_cache.common import (
kv_to_page_indices,
page_align_floor,
release_kv_cache,
retraction_discard,
retraction_restore,
)
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.mem_cache.memory_pool import (
@@ -693,6 +695,12 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
def release_memory_occupation(self):
self.queue.clear()
for req in self.retracted_queue:
retraction_discard(
req,
self.tree_cache,
get_disagg().disaggregation_decode_retraction_backup,
)
self.retracted_queue.clear()
if hasattr(self.kv_manager, "deregister_buffer_to_engine"):
self.kv_manager.deregister_buffer_to_engine()
@@ -740,8 +748,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
if uses_swa_tail_prealloc:
swa_allocatable_tokens -= swa_required
# load from cpu, release the cpu copy
req.load_kv_cache(self.req_to_token_pool, self.token_to_kv_pool_allocator)
retraction_restore(
req,
self.tree_cache,
self.req_to_token_pool,
self.token_to_kv_pool_allocator,
get_disagg().disaggregation_decode_retraction_backup,
)
self.retracted_queue = [
entry
+29 -7
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
from sglang.srt.dllm.config import DllmConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_schedule,
get_serving,
@@ -97,9 +98,11 @@ from sglang.srt.mem_cache.base_prefix_cache import (
zero_match_result,
)
from sglang.srt.mem_cache.common import (
RetractionBackup,
evict_from_tree_cache,
free_swa_out_of_window_slots,
release_kv_cache,
retraction_backup,
)
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey
@@ -878,6 +881,7 @@ class Req(ReqDllmMixin):
# For req-level memory management
self.kv_committed_len = 0
self.kv: Optional[ReqKvInfo] = None
self.retraction_backup: Optional[RetractionBackup] = None
# for cross-encoder model
self.token_type_ids = token_type_ids
@@ -1712,19 +1716,24 @@ class Req(ReqDllmMixin):
self.req_pool_idx, : self.seqlen - 1
]
# Copies over both the kv cache and mamba state if available
self.kv_cache_cpu = token_to_kv_pool_allocator.get_cpu_copy(
token_indices, mamba_indices=self.mamba_pool_idx
self.retraction_backup = RetractionBackup(
cpu_tensors=token_to_kv_pool_allocator.get_cpu_copy(
token_indices, mamba_indices=self.mamba_pool_idx
)
)
def load_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator):
assert self.retraction_backup is not None
token_indices = req_to_token_pool.req_to_token[
self.req_pool_idx, : self.seqlen - 1
]
# Loads both the kv cache and mamba state if exists
token_to_kv_pool_allocator.load_cpu_copy(
self.kv_cache_cpu, token_indices, mamba_indices=self.mamba_pool_idx
self.retraction_backup.cpu_tensors,
token_indices,
mamba_indices=self.mamba_pool_idx,
)
del self.kv_cache_cpu
self.retraction_backup = None
def build_rebootstrap_payload(self) -> dict:
"""Build the prefill ``/generate`` payload that asks the original prefill
@@ -1914,7 +1923,13 @@ def release_req(
# Callers that will recompute the KV instead (PD true-retraction rebootstrap)
# pass offload_kv=False to skip the wasteful device->host copy.
if server_args.disaggregation_mode == "decode" and offload_kv:
req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
retraction_backup(
req,
tree_cache,
req_to_token_pool,
token_to_kv_pool_allocator,
get_disagg().disaggregation_decode_retraction_backup,
)
# TODO (csy): for preempted requests, we may want to insert into the tree
release_kv_cache(req, tree_cache, is_insert=False)
# NOTE(lsyin): we should use the newly evictable memory instantly.
@@ -2834,7 +2849,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
)
reqs_to_abort.append(last_req)
self.release_req(last_idx, 0, server_args)
self.release_req(last_idx, 0, server_args, offload_kv=False)
logger.warning(
"retract_decode: aborted last request %s due to OOM", last_req.rid
)
@@ -2889,7 +2904,13 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
)
return sorted_indices
def release_req(self, idx: int, remaing_req_count: int, server_args: ServerArgs):
def release_req(
self,
idx: int,
remaing_req_count: int,
server_args: ServerArgs,
offload_kv: bool = True,
):
release_req(
req=self.reqs[idx],
remaing_req_count=remaing_req_count,
@@ -2898,6 +2919,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator,
offload_kv=offload_kv,
)
def prepare_encoder_info_decode(self):
+14 -4
View File
@@ -266,7 +266,11 @@ from sglang.srt.managers.utils import (
validate_input_length,
)
from sglang.srt.mem_cache import kv_cache_builder
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req,
release_kv_cache,
retraction_discard,
)
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
from sglang.srt.model_loader.utils import get_resolved_model_impl
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
@@ -965,6 +969,9 @@ class Scheduler(
def init_memory_pools(self):
"""Allocate KV cache pools for target and draft workers."""
self.init_target_memory_pool()
# Lands the retraction backend on the disagg bag before the draft
# worker's HiCache plan reads it.
kv_cache_builder.resolve_decode_retraction_backup(tp_worker=self.tp_worker)
if self.draft_worker is not None:
pool, allocator = self.tp_worker.get_memory_pool()
self.draft_worker.alloc_memory_pool(
@@ -4540,13 +4547,16 @@ class Scheduler(
logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}")
decode_req.kv_receiver.abort()
# Abort requests already retracted to CPU cache
# Abort requests whose KV is already backed up for retraction.
if self.disagg_decode_prealloc_queue.retracted_queue:
remaining_retracted = []
for decode_req in self.disagg_decode_prealloc_queue.retracted_queue:
if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid):
assert hasattr(decode_req, "kv_cache_cpu")
del decode_req.kv_cache_cpu
retraction_discard(
decode_req,
self.tree_cache,
get_disagg().disaggregation_decode_retraction_backup,
)
self.ipc_channels.send_to_tokenizer.send_output(
AbortReq(rid=decode_req.rid), decode_req
)
+63 -1
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, NamedTuple, Optional, cast
import numpy as np
import torch
@@ -12,6 +12,7 @@ from sglang.kernels.ops.memory.common import (
from sglang.kernels.ops.memory.common import get_last_loc_kernel as get_last_loc_kernel
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.hicache_storage import PoolTransfer
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_serving, get_spec
from sglang.srt.utils.common import ceil_align
@@ -19,6 +20,7 @@ from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
# Needs 2 + 1 slots for mamba request with prefix cache. 2 for ping pong cache, 1 for running mamba state.
MAMBA_STATE_PER_REQ_PREFIX_CACHE = 3
@@ -29,6 +31,12 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__)
class RetractionBackup(NamedTuple):
cpu_tensors: Any = None
host_indices: Optional[torch.Tensor] = None
pool_transfers: Optional[list[PoolTransfer]] = None
def kv_to_page_indices(kv_indices: torch.Tensor, page_size: int) -> np.ndarray:
return (kv_indices[::page_size] // page_size).cpu().numpy()
@@ -130,6 +138,60 @@ def evict_from_tree_cache(tree_cache: BasePrefixCache | None, num_tokens: int):
tree_cache.evict(EvictParams(num_tokens=num_tokens - available_size))
def retraction_backup(
req: Req,
tree_cache: BasePrefixCache,
req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
backend: str,
) -> None:
if backend == "cpu_tensor":
req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
return
if backend != "host_pool":
raise ValueError(f"Unknown retraction backup backend: {backend}")
if req.seqlen <= 1:
return
unified_cache = cast("UnifiedRadixCache", tree_cache)
req.retraction_backup = unified_cache.retraction_backup(req)
def retraction_restore(
req: Req,
tree_cache: BasePrefixCache,
req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
backend: str,
) -> None:
if backend == "cpu_tensor":
req.load_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
return
if backend != "host_pool":
raise ValueError(f"Unknown retraction backup backend: {backend}")
if req.seqlen <= 1:
return
unified_cache = cast("UnifiedRadixCache", tree_cache)
assert req.retraction_backup is not None
unified_cache.retraction_restore(req, req.retraction_backup)
req.retraction_backup = None
def retraction_discard(req: Req, tree_cache: BasePrefixCache, backend: str) -> None:
if backend == "cpu_tensor":
req.retraction_backup = None
return
if backend != "host_pool":
raise ValueError(f"Unknown retraction backup backend: {backend}")
if req.retraction_backup is None:
return
unified_cache = cast("UnifiedRadixCache", tree_cache)
unified_cache.retraction_discard(req.retraction_backup)
req.retraction_backup = None
def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = True):
# the two resources currently have the same lifecycle, thus simplify logic below
assert (req.req_pool_idx is None) == (req.kv is None)
+3 -2
View File
@@ -65,6 +65,7 @@ from sglang.srt.observability.metrics_collector import (
StorageMetricsCollector,
resolve_collector_class,
)
from sglang.srt.runtime_context import get_memory
if TYPE_CHECKING:
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
@@ -86,7 +87,7 @@ class HiRadixCache(RadixCache):
if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
self.kv_cache,
server_args.hicache_ratio,
get_memory().hicache_ratio,
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
@@ -104,7 +105,7 @@ class HiRadixCache(RadixCache):
_parallel = get_parallel()
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache,
server_args.hicache_ratio,
get_memory().hicache_ratio,
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
@@ -28,7 +28,7 @@ from sglang.srt.mem_cache.pool_host.mha import (
)
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import get_memory, get_parallel
if TYPE_CHECKING:
import torch
@@ -99,7 +99,7 @@ def build_kv_host_pool(
kwargs["dcp_rank"] = parallel.attn_dcp_rank
return kv_host_pool_cls(
kv_pool,
server_args.hicache_ratio,
get_memory().hicache_ratio,
server_args.hicache_size if host_size is None else host_size,
page_size,
server_args.hicache_mem_layout,
@@ -402,7 +402,7 @@ def _deepseek_v4_num_host_pages(
"DeepSeek V4 HiCache currently does not support --hicache-size; "
"use --hicache-ratio instead."
)
ratio = server_args.hicache_ratio
ratio = get_memory().hicache_ratio
full_host_pages = int(device_full_pages * ratio)
swa_host_pages = int(device_swa_pages * ratio)
return full_host_pages, swa_host_pages
@@ -715,7 +715,7 @@ def build_hybrid_mamba_stack(
)
mamba_host_pool = MambaPoolHost(
mamba_pool,
server_args.hicache_ratio,
get_memory().hicache_ratio,
mamba_host_size,
allocator_type=_get_allocator_type(server_args),
layout=server_args.hicache_mem_layout,
@@ -819,7 +819,7 @@ def build_hybrid_mamba_swa_stack(
)
mamba_host_pool = MambaPoolHost(
mamba_pool,
server_args.hicache_ratio,
get_memory().hicache_ratio,
mamba_host_size,
allocator_type=server_args.hicache_storage_backend,
layout=server_args.hicache_mem_layout,
@@ -34,9 +34,18 @@ from sglang.srt.configs.model_config import ModelImpl, is_deepseek_dsa
from sglang.srt.environ import envs
from sglang.srt.managers.mm_schedule import init_mm_embedding_cache
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
from sglang.srt.mem_cache.registry import TreeCacheBuildContext, create_tree_cache
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.model_loader.utils import get_resolved_model_impl
from sglang.srt.runtime_context import get_parallel, get_schedule
from sglang.srt.runtime_context import (
get_context,
get_disagg,
get_memory,
get_parallel,
get_schedule,
)
if TYPE_CHECKING:
@@ -130,6 +139,57 @@ def _register_legacy_hicache_draft(
tree_cache.cache_controller.set_draft_kv_pool(pool, draft_host_pool)
def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str:
"""Resolve the retraction backend onto the config bags and return it.
The backend needs the built KV pool, so it cannot resolve in
``ServerArgs.__post_init__``; it lands on the bags via ``override`` and
every reader goes through ``get_disagg()`` / ``get_memory()``.
"""
disagg = get_disagg()
memory = get_memory()
fields = {}
backend = disagg.disaggregation_decode_retraction_backup
if backend is None:
kv_cache = tp_worker.get_memory_pool()[1].get_kvcache()
full_tokens_per_layer = (
tp_worker.get_tokens_per_layer_info()[0]
if tp_worker.is_hybrid_swa
else None
)
supports_host_pool = isinstance(kv_cache, MHATokenToKVPool) or (
isinstance(kv_cache, SWAKVPool) and full_tokens_per_layer > 0
)
schedule = get_schedule()
priority_preemption = (
schedule.enable_priority_scheduling
and not schedule.disable_priority_preemption
)
backend = (
"host_pool"
if disagg.disaggregation_mode == "decode"
and not get_parallel().dcp_enabled
and not disagg.disaggregation_decode_enable_radix_cache
# KV offload already owns a host pool; a second one double-books host memory.
and not disagg.disaggregation_decode_enable_offload_kvcache
and not priority_preemption
and supports_host_pool
else "cpu_tensor"
)
fields["disaggregation_decode_retraction_backup"] = backend
if memory.hicache_ratio is None:
# Only a decode server reaches resolution with the ratio unset; host-pool
# retraction sizes the host pool 1:1 with the device pool, everything
# else keeps the standard default.
fields["hicache_ratio"] = 1.0 if backend == "host_pool" else 2.0
source = "kv_cache_builder.decode_retraction"
get_context().override(source, **fields)
return backend
def build_kv_cache(
*,
server_args: ServerArgs,
@@ -178,6 +238,8 @@ def build_kv_cache(
req_to_token_pool, token_to_kv_pool_allocator = tp_worker.get_memory_pool()
mtp_draft_device_pools = tp_worker.model_runner.mtp_draft_device_pools
retraction_backup = resolve_decode_retraction_backup(tp_worker=tp_worker)
disable_radix_cache = server_args.disable_radix_cache or (
model_config.is_multimodal and uses_transformers_backend
)
@@ -260,7 +322,9 @@ def build_kv_cache(
)
)
if enable_hierarchical_cache and hicache_draft_plan is not None:
if (
enable_hierarchical_cache or retraction_backup == "host_pool"
) and hicache_draft_plan is not None:
maybe_register_hicache_draft(
tree_cache=tree_cache,
draft_plan=hicache_draft_plan,
@@ -268,6 +332,14 @@ def build_kv_cache(
page_size=page_size,
)
if retraction_backup == "host_pool":
if not isinstance(tree_cache, UnifiedRadixCache):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool requires "
"UnifiedRadixCache with HiCache attached."
)
tree_cache.validate_retraction_host_capacity()
embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024)
+20 -4
View File
@@ -17,7 +17,7 @@ from typing import TYPE_CHECKING, Any, Callable, Optional
from sglang.srt.environ import envs
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.runtime_context import get_memory
from sglang.srt.runtime_context import get_disagg, get_memory
from sglang.srt.utils.tensor_bridge import use_mlx
if TYPE_CHECKING:
@@ -82,6 +82,12 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
server_args = ctx.server_args
params = ctx.params
if (
ctx.disable_radix_cache
and get_disagg().disaggregation_decode_retraction_backup == "host_pool"
):
return _create_unified_radix_cache(ctx, server_args, params)
if ctx.effective_chunked_prefill_size is not None and ctx.disable_radix_cache:
if not ctx.is_hybrid_swa:
from sglang.srt.mem_cache.chunk_cache import ChunkCache
@@ -167,6 +173,12 @@ def _create_unified_radix_cache(
params: CacheInitParams,
) -> BasePrefixCache:
"""Initialize a UnifiedRadixCache with proper components and optional HiCache."""
if get_disagg().disaggregation_decode_retraction_backup == "host_pool":
if ctx.is_hybrid_ssm:
raise ValueError("Host-pool retraction does not support Mamba models.")
if ctx.is_hybrid_swa and ctx.full_tokens_per_layer == 0:
raise ValueError("Host-pool retraction does not support pure-SWA models.")
from sglang.srt.mem_cache.unified_cache.components import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
@@ -197,7 +209,10 @@ def _create_unified_radix_cache(
ComponentType.MAMBA: MlxAuxiliaryStateComponent,
}
cache = UnifiedRadixCache(params)
if ctx.enable_hierarchical_cache:
if (
ctx.enable_hierarchical_cache
or get_disagg().disaggregation_decode_retraction_backup == "host_pool"
):
cache.init_hicache(server_args, params)
ctx.tp_worker.register_hicache_layer_transfer_counter(
cache.cache_controller.layer_done_counter
@@ -232,6 +247,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
"--enable-session-radix-cache)."
)
hicache_attached = cache.cache_controller is not None
streaming_wrapped = False
if (
ctx.server_args.enable_streaming_session
@@ -244,12 +260,12 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
logger.info(
"Tree cache initialized: source=%s impl=%s hybrid_swa=%s hybrid_ssm=%s "
"hierarchical=%s streaming_wrapped=%s",
"hicache_attached=%s streaming_wrapped=%s",
source,
type(cache).__name__,
ctx.is_hybrid_swa,
ctx.is_hybrid_ssm,
ctx.enable_hierarchical_cache,
hicache_attached,
streaming_wrapped,
)
return cache
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging
import threading
import time
from dataclasses import replace
from queue import Empty, Queue
from typing import TYPE_CHECKING, Iterator, NamedTuple, Optional, Sequence, TypeVar
@@ -10,6 +11,7 @@ import torch
from sglang.srt.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs
from sglang.srt.managers.cache_controller import CacheOperation
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
DecLockRefParams,
@@ -23,6 +25,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.common import RetractionBackup
from sglang.srt.mem_cache.hicache_storage import (
PoolHitPolicy,
PoolName,
@@ -32,7 +35,9 @@ from sglang.srt.mem_cache.hicache_storage import (
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
HybridCacheController,
)
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
from sglang.srt.mem_cache.unified_cache.cache_action import (
BackupKV,
CacheAction,
@@ -70,6 +75,7 @@ from sglang.srt.observability.metrics_collector import (
resolve_collector_class,
)
from sglang.srt.session.streaming_session import StreamingSession
from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING:
from sglang.srt.managers.cache_controller import HiCacheAck
@@ -425,6 +431,9 @@ class UnifiedRadixCache(BasePrefixCache):
assert not result.cache_actions
return result
def is_chunk_cache(self) -> bool:
return self.disable
def insert(self, params: InsertParams) -> InsertResult:
if self.disable:
return InsertResult(prefix_len=0)
@@ -921,6 +930,240 @@ class UnifiedRadixCache(BasePrefixCache):
self._free_values(result.device_frees, result.host_frees)
return result.tracker.get(component_type, 0)
# ---- Decode retraction ----
def supports_retraction_backup(self) -> bool:
if self.cache_controller is None or self.host_pool_group is None:
return False
if self.supports_mamba():
return False
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
if isinstance(kv_cache, SWAKVPool):
return (
self.supports_swa()
and {
PoolName.KV,
PoolName.SWA,
}
<= self.host_pool_group.entry_map.keys()
)
return isinstance(kv_cache, MHATokenToKVPool) and (
PoolName.KV in self.host_pool_group.entry_map
)
def validate_retraction_host_capacity(self) -> None:
if not self.supports_retraction_backup():
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool requires "
"an MHA or hybrid-SWA HiCache host stack."
)
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
device_pools = {PoolName.KV: kv_cache}
if isinstance(kv_cache, SWAKVPool):
device_pools = {
PoolName.KV: kv_cache.full_kv_pool,
PoolName.SWA: kv_cache.swa_kv_pool,
}
for name, device_pool in device_pools.items():
host_pool = self.host_pool_group.entry_map[name].host_pool
if host_pool.logical_size < device_pool.size:
raise ValueError(
"Retraction host pool is smaller than its device pool: "
f"pool={name}, host_slots={host_pool.logical_size}, "
f"device_slots={device_pool.size}. Increase --hicache-ratio "
"or --hicache-size."
)
for spec in self.sidecar_pool_specs:
source_size = self.host_pool_group.entry_map[
spec.indices_from_pool
].host_pool.logical_size
sidecar_size = self.host_pool_group.entry_map[
spec.pool_name
].host_pool.logical_size
if sidecar_size < source_size:
raise ValueError(
"Retraction sidecar host pool is smaller than its index source: "
f"pool={spec.pool_name}, host_slots={sidecar_size}, "
f"source={spec.indices_from_pool}, source_slots={source_size}."
)
@staticmethod
def _pad_retraction_indices(indices: torch.Tensor, page_size: int) -> torch.Tensor:
aligned_len = ceil_align(len(indices), page_size)
if aligned_len == len(indices):
return indices
tail = indices[-1] + torch.arange(
1,
aligned_len - len(indices) + 1,
dtype=torch.int64,
device=indices.device,
)
return torch.cat([indices, tail])
def _retraction_device_transfers(
self, req: Req
) -> tuple[torch.Tensor, list[PoolTransfer]]:
num_tokens = req.seqlen - 1
full_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, :num_tokens
].to(torch.int64)
full_indices = self._pad_retraction_indices(full_indices, self.page_size)
component_transfers: dict[ComponentType, list[PoolTransfer]] = {}
if self.supports_swa():
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
assert self.sliding_window_size is not None
window_start = max(0, num_tokens - self.sliding_window_size)
window_start = window_start // self.page_size * self.page_size
window_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, window_start:num_tokens
].to(torch.int64)
swa_indices = kv_cache.translate_loc_from_full_to_swa(window_indices)
assert bool(
(swa_indices > 0).all()
), f"unmapped SWA window positions for request {req.rid}"
component_transfers[ComponentType.SWA] = [
PoolTransfer(
name=PoolName.SWA,
device_indices=self._pad_retraction_indices(
swa_indices, self.page_size
),
)
]
kv_transfer = PoolTransfer(name=PoolName.KV, device_indices=full_indices)
extra_transfers = [
transfer
for transfers in component_transfers.values()
for transfer in transfers
]
extra_transfers.extend(
self._build_sidecar_transfers(
CacheTransferPhase.BACKUP_HOST,
kv_transfer,
component_transfers,
)
)
return full_indices, extra_transfers
def _reclaim_retraction_host(self, num_tokens: int) -> int:
if self.disable:
return 0
return self.evict_host(num_tokens)
def retraction_backup(self, req: Req) -> RetractionBackup:
assert req.seqlen > 1
device_indices, extra_transfers = self._retraction_device_transfers(req)
host_indices = self.host_pool_group.alloc(len(device_indices))
if host_indices is None:
self._reclaim_retraction_host(len(device_indices))
host_indices = self.host_pool_group.alloc(len(device_indices))
if host_indices is None:
raise RuntimeError(
"Retraction host KV pool exhausted after reclaim: "
f"request={req.rid}, required_slots={len(device_indices)}, "
f"available_slots={self.host_pool_group.available_size()}."
)
resolved = self.cache_controller._resolve_pool_transfers_allocation(
extra_transfers or None,
alloc_host=True,
kv_device_indices=device_indices,
kv_host_indices=host_indices,
)
if resolved is None and extra_transfers:
self.host_pool_group.free(host_indices)
raise RuntimeError(
"Retraction auxiliary host allocation failed after atomic rollback: "
f"request={req.rid}, pools={[x.name for x in extra_transfers]}."
)
backup = RetractionBackup(
host_indices=host_indices,
pool_transfers=[replace(x, device_indices=None) for x in resolved or []]
or None,
)
operation = CacheOperation(
host_indices,
device_indices,
node_id=-1,
pool_transfers=resolved,
)
try:
write_host, write_device, write_pools = (
self.cache_controller._move_write_operation(operation)
)
completion = self.cache_controller.l2_transfer_engine.submit_device_to_host(
self.cache_controller._l2_transfers(
write_host, write_device, write_pools
)
)
completion.finish_event.synchronize()
except Exception:
self.retraction_discard(backup)
raise
return backup
def retraction_restore(self, req: Req, backup: RetractionBackup) -> None:
device_indices, current_transfers = self._retraction_device_transfers(req)
assert len(backup.host_indices) == len(device_indices), (
f"Host backup has {len(backup.host_indices)} slots, but restore has "
f"{len(device_indices)}"
)
current_by_name = {transfer.name: transfer for transfer in current_transfers}
saved_by_name = {
transfer.name: transfer for transfer in backup.pool_transfers or []
}
assert current_by_name.keys() == saved_by_name.keys(), (
f"Host backup pools {set(saved_by_name)} do not match restore pools "
f"{set(current_by_name)}"
)
restored_transfers = [
replace(
saved,
device_indices=current_by_name[name].device_indices,
)
for name, saved in saved_by_name.items()
]
resolved = self.cache_controller._resolve_pool_transfers_allocation(
restored_transfers or None,
alloc_host=False,
kv_device_indices=device_indices,
kv_host_indices=backup.host_indices,
)
assert resolved is not None or not restored_transfers
operation = CacheOperation(
backup.host_indices,
device_indices,
node_id=-1,
pool_transfers=resolved,
)
load_host, load_device, load_pools = self.cache_controller._move_op_indices(
operation
)
completion = self.cache_controller.l2_transfer_engine.submit_host_to_device(
self.cache_controller._l2_load_transfers(
load_host, load_device, load_pools
),
layer_num=self.cache_controller.layer_num,
)
completion.finish_event.synchronize()
self.retraction_discard(backup)
def retraction_discard(self, backup: RetractionBackup) -> None:
self.host_pool_group.free(backup.host_indices)
for transfer in backup.pool_transfers or []:
if transfer.indices_from_pool is None:
assert transfer.host_indices is not None
self.host_pool_group.get_pool(transfer.name).free(transfer.host_indices)
# ---- HiCache: Backup / LoadBack ----
def _execute_and_commit_kv_backup(
+57 -3
View File
@@ -2653,10 +2653,10 @@ class ServerArgs:
False
)
hicache_ratio: A[
float,
"The ratio of the size of host KV cache memory pool to the size of device pool.",
Optional[float],
"The ratio of the size of host KV cache memory pool to the size of device pool. Defaults to 2.0, or 1.0 for host-pool decode retraction.",
NS("memory"),
] = 2.0
] = None
hicache_size: A[
int,
"The size of host KV cache memory pool in gigabytes, which will override the hicache_ratio if set.",
@@ -3115,6 +3115,19 @@ class ServerArgs:
"Enable async KV cache offloading on decode server (PD mode).",
NS("disagg"),
] = False
disaggregation_decode_retraction_backup: A[
Optional[str],
Arg(
help=(
"Storage backend for KV preserved across PD decode retraction. "
"'cpu_tensor' uses per-request CPU tensors. 'host_pool' uses "
"a reserved HiCache pool and does not fall back on exhaustion. "
"If omitted, the backend is inferred from the decode KV pool."
),
choices=["cpu_tensor", "host_pool"],
),
NS("disagg"),
] = None
num_reserved_decode_tokens: A[
int,
"Number of decode tokens that will have memory reserved when adding new request to the running batch.",
@@ -3589,6 +3602,7 @@ class ServerArgs:
self._handle_moe_runner_backend_alias()
self._handle_return_hidden_states_mode()
self._handle_media_url_security()
self._handle_hicache_ratio_default()
if self.model_path.lower() in ["none", "dummy"]:
return
@@ -7335,6 +7349,10 @@ class ServerArgs:
"workloads and backends may be supported in a future change."
)
def _handle_hicache_ratio_default(self):
if self.hicache_ratio is None and self.disaggregation_mode != "decode":
self.hicache_ratio = 2.0
def _handle_hicache(self):
"""Normalize hicache-related knobs into a valid runtime configuration.
@@ -7346,6 +7364,10 @@ class ServerArgs:
if not (
self.enable_hierarchical_cache
or self.disaggregation_decode_enable_offload_kvcache
or (
self.disaggregation_mode == "decode"
and self.disaggregation_decode_retraction_backup in (None, "host_pool")
)
):
return
@@ -8084,6 +8106,32 @@ class ServerArgs:
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
def _handle_cache_compatibility(self):
if (
self.disaggregation_decode_retraction_backup == "host_pool"
and self.disaggregation_mode != "decode"
):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool is only "
"supported on a PD decode server."
)
if (
self.disaggregation_decode_retraction_backup == "host_pool"
and self.dcp_size > 1
):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool does not "
"support --dcp-size > 1."
)
if (
self.disaggregation_decode_retraction_backup == "host_pool"
and self.enable_priority_scheduling
and not self.disable_priority_preemption
):
raise ValueError(
"--disaggregation-decode-retraction-backup=host_pool requires "
"--disable-priority-preemption when priority scheduling is enabled."
)
if self.enable_hierarchical_cache and self.disable_radix_cache:
raise ValueError(
"The arguments enable-hierarchical-cache and disable-radix-cache are mutually exclusive "
@@ -8099,6 +8147,12 @@ class ServerArgs:
raise ValueError(
"The argument disaggregation-decode-enable-offload-kvcache is only supported when hicache-storage-backend is provided."
)
if self.disaggregation_decode_retraction_backup == "host_pool":
raise ValueError(
"The arguments disaggregation-decode-enable-offload-kvcache and "
"disaggregation-decode-retraction-backup=host_pool are mutually exclusive: "
"both build a decode host pool."
)
# Validate the effective ratio: model branches may declare a reset
# (e.g. Step3p forces 1.0 under hierarchical cache) that supersedes
@@ -11,7 +11,7 @@ from sglang.srt.model_executor.graph_memory_usage import (
merge_graph_memory_usage,
merge_graph_time_usage,
)
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule
from sglang.srt.runtime_context import get_disagg, get_exec, get_memory, get_schedule
if TYPE_CHECKING:
from sglang.srt.managers.io_struct import (
@@ -235,7 +235,10 @@ class BaseSpecWorker(ABC):
target_model_runner = self.target_worker.model_runner
target_model_runner.mtp_draft_device_pools = ()
spec_algorithm = target_model_runner.spec_algorithm
if not get_memory().enable_hierarchical_cache:
if not (
get_memory().enable_hierarchical_cache
or get_disagg().disaggregation_decode_retraction_backup == "host_pool"
):
return HiCacheDraftPlan()
draft_runners = self._draft_model_runners()
@@ -225,13 +225,15 @@ class TestDisaggregationMooncakeFailure(PDDisaggregationServerBase):
class TestDisaggregationMooncakeSpec(
JSONConstrainedMixin, SpecGrammarKit, PDDisaggregationServerBase
):
min_retraction_accept_length = 1.3
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
spec_args = [
"--speculative-algorithm",
"EAGLE",
"EAGLE3",
"--speculative-draft-model-path",
DEFAULT_DRAFT_MODEL_EAGLE3,
"--speculative-num-steps",
@@ -245,9 +247,51 @@ class TestDisaggregationMooncakeSpec(
"--dtype=float16",
]
cls.extra_prefill_args = spec_args
cls.extra_decode_args = spec_args
cls.extra_decode_args = [
*spec_args,
"--disaggregation-decode-retraction-backup",
"host_pool",
]
cls.extra_decode_env = {"SGLANG_TEST_RETRACT": "true"}
cls.launch_all()
def test_host_pool_retraction_preserves_spec_acceptance(self):
prompts = [
f"Request {i}: explain how speculative decoding works. " * 4
for i in range(4)
]
response = requests.post(
self.lb_url + "/generate",
json={
"text": prompts,
"sampling_params": {
"temperature": 0,
"ignore_eos": True,
"max_new_tokens": 64,
},
},
)
response.raise_for_status()
results = response.json()
retracted_results = [
result for result in results if result["meta_info"]["num_retractions"] > 0
]
retraction_count = sum(
result["meta_info"]["num_retractions"] for result in retracted_results
)
self.assertGreater(retraction_count, 0)
completion_tokens = sum(
result["meta_info"]["completion_tokens"] for result in retracted_results
)
verify_count = sum(
result["meta_info"]["spec_verify_ct"] for result in retracted_results
)
self.assertGreater(verify_count, 0)
accept_length = completion_tokens / verify_count
print(f"Retraction speculative {accept_length=:.4f}")
self.assertGreater(accept_length, self.min_retraction_accept_length)
def test_gsm8k(self):
args = SimpleNamespace(
base_url=f"http://{self.base_host}:{self.lb_port}",
@@ -12,6 +12,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -31,6 +32,14 @@ class FakeReceiver:
class TestDecodeQueueCleanup(CustomTestCase):
def test_paged_swa_retraction_resume_uses_physical_page_budget(self):
# resume_retracted_reqs reads the retraction backend off the disagg
# bag, so the case publishes a config instead of injecting one.
override = get_context().override_server_args(
disaggregation_decode_retraction_backup="cpu_tensor"
)
override.install()
self.addCleanup(override.restore)
page_size = 128
fill_len = 574
physical_tokens_per_req = 5 * page_size
@@ -42,6 +51,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
origin_input_ids=[0] * fill_len,
output_ids=[],
is_retracted=True,
retraction_backup=None,
load_kv_cache=MagicMock(),
)
for i in range(4)
@@ -52,6 +62,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue.num_reserved_decode_tokens = 0
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
queue.tree_cache = MagicMock()
queue.scheduler = SimpleNamespace(
sliding_window_size=2047,
server_args=SimpleNamespace(disable_radix_cache=True),
@@ -0,0 +1,173 @@
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.hicache_storage import PoolName
from sglang.srt.mem_cache.kv_cache_builder import maybe_register_hicache_draft
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
from sglang.srt.mem_cache.unified_cache.components import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.speculative.base_spec_worker import (
HiCacheDraftMode,
HiCacheDraftPlan,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
class TestDecodeRetractionBackup(unittest.TestCase):
pool_size = 32
num_tokens = 8
dtype = torch.bfloat16
device = "cuda"
def _make_pool(self, layer_num: int) -> MHATokenToKVPool:
return MHATokenToKVPool(
size=self.pool_size,
page_size=1,
head_num=2,
head_dim=64,
dtype=self.dtype,
layer_num=layer_num,
device=self.device,
enable_memory_saver=False,
)
def _seed_pool(
self, pool: MHATokenToKVPool, indices: torch.Tensor, base: int
) -> None:
for layer_id, (key, value) in enumerate(
zip(pool.k_buffer, pool.v_buffer, strict=True)
):
pattern = torch.arange(
key[indices].numel(), device=self.device, dtype=torch.float32
).reshape_as(key[indices])
key[indices] = (pattern + base + layer_id * 100).to(self.dtype)
value[indices] = (pattern + base + 50 + layer_id * 100).to(self.dtype)
@staticmethod
def _snapshot_pool(
pool: MHATokenToKVPool, indices: torch.Tensor
) -> list[tuple[torch.Tensor, torch.Tensor]]:
return [
(key[indices].clone(), value[indices].clone())
for key, value in zip(pool.k_buffer, pool.v_buffer, strict=True)
]
def _assert_pool_equal(
self,
pool: MHATokenToKVPool,
indices: torch.Tensor,
expected: list[tuple[torch.Tensor, torch.Tensor]],
) -> None:
for (key, value), (expected_key, expected_value) in zip(
zip(pool.k_buffer, pool.v_buffer, strict=True), expected, strict=True
):
self.assertTrue(torch.equal(key[indices], expected_key))
self.assertTrue(torch.equal(value[indices], expected_value))
def test_restores_target_and_draft_kv(self):
server_args = ServerArgs(
model_path="dummy",
page_size=1,
hicache_ratio=1.0,
hicache_io_backend="kernel",
hicache_mem_layout="page_first",
)
set_global_server_args_for_scheduler(server_args)
req_to_token_pool = ReqToTokenPool(
size=2,
max_context_len=self.pool_size,
device=self.device,
enable_memory_saver=False,
)
target_pool = self._make_pool(layer_num=2)
allocator = TokenToKVPoolAllocator(
size=self.pool_size,
dtype=self.dtype,
device=self.device,
kvcache=target_pool,
need_sort=False,
)
params = CacheInitParams(
disable=True,
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=1,
is_eagle=True,
tree_components=(ComponentType.FULL,),
)
cache = UnifiedRadixCache(params)
cache.init_hicache(server_args, params)
self.addCleanup(cache.release_host_resources)
draft_pool = self._make_pool(layer_num=1)
maybe_register_hicache_draft(
tree_cache=cache,
draft_plan=HiCacheDraftPlan(
mode=HiCacheDraftMode.SIDECAR,
device_pools=(draft_pool,),
),
server_args=server_args,
page_size=1,
)
self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map)
cache.validate_retraction_host_capacity()
req = SimpleNamespace(
rid="request", req_pool_idx=None, seqlen=self.num_tokens + 1
)
self.assertIsNotNone(req_to_token_pool.alloc([req]))
source_indices = allocator.alloc(self.num_tokens)
self.assertIsNotNone(source_indices)
req_to_token_pool.write(
(req.req_pool_idx, slice(0, self.num_tokens)), source_indices
)
self._seed_pool(target_pool, source_indices, base=1000)
self._seed_pool(draft_pool, source_indices, base=3000)
target_expected = self._snapshot_pool(target_pool, source_indices)
draft_expected = self._snapshot_pool(draft_pool, source_indices)
host_free_before = cache.host_pool_group.available_size()
backup = cache.retraction_backup(req)
self.assertEqual(
{transfer.name for transfer in backup.pool_transfers or []},
{PoolName.DRAFT},
)
self.assertLess(cache.host_pool_group.available_size(), host_free_before)
for buffer in (*target_pool.k_buffer, *target_pool.v_buffer):
buffer.fill_(-1)
for buffer in (*draft_pool.k_buffer, *draft_pool.v_buffer):
buffer.fill_(-2)
allocator.free(source_indices)
blocker_indices = allocator.alloc(self.num_tokens)
destination_indices = allocator.alloc(self.num_tokens)
self.assertIsNotNone(blocker_indices)
self.assertIsNotNone(destination_indices)
self.assertFalse(torch.equal(source_indices, destination_indices))
req_to_token_pool.write(
(req.req_pool_idx, slice(0, self.num_tokens)), destination_indices
)
cache.retraction_restore(req, backup)
self._assert_pool_equal(target_pool, destination_indices, target_expected)
self._assert_pool_equal(draft_pool, destination_indices, draft_expected)
self.assertEqual(cache.host_pool_group.available_size(), host_free_before)
allocator.free(blocker_indices)
allocator.free(destination_indices)
req_to_token_pool.free(req)
if __name__ == "__main__":
unittest.main()
@@ -1331,6 +1331,27 @@ class TestHiCacheArgs(unittest.TestCase):
self.assertEqual(args.hicache_mem_layout, "page_first")
self.assertIsNone(args.decode_attention_backend)
def test_decode_offload_rejects_host_pool_retraction(self):
args = self._make_args(
disaggregation_mode="decode",
disaggregation_decode_enable_offload_kvcache=True,
hicache_storage_backend="file",
disaggregation_decode_retraction_backup="host_pool",
)
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
args._handle_cache_compatibility()
def test_decode_offload_allows_cpu_tensor_retraction(self):
args = self._make_args(
disaggregation_mode="decode",
disaggregation_decode_enable_offload_kvcache=True,
hicache_storage_backend="file",
disaggregation_decode_retraction_backup="cpu_tensor",
)
args._handle_cache_compatibility()
class TestNgramExternalSamArgs(CustomTestCase):
def _make_dummy_ngram_args(self, **overrides):