diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index c26de286c..56febf102 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index e927daeff..3ff2a29f3 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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): diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4bac671ad..53a60dc10 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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 ) diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 88f6a8e30..4bc61811b 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -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) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 1b9111608..cee15cb05 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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, diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 0f1e194e0..bedf086a6 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -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, diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index a1459d982..3a6554bb5 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -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) diff --git a/python/sglang/srt/mem_cache/registry.py b/python/sglang/srt/mem_cache/registry.py index e979e08ad..8aaec2641 100644 --- a/python/sglang/srt/mem_cache/registry.py +++ b/python/sglang/srt/mem_cache/registry.py @@ -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 diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 1ca6e6583..3db4fa36e 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 078650a37..3601485ad 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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 diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index 8111c9658..ba0ffd616 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -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() diff --git a/test/registered/disaggregation/test_disaggregation_basic.py b/test/registered/disaggregation/test_disaggregation_basic.py index f010c0416..0e1c75230 100644 --- a/test/registered/disaggregation/test_disaggregation_basic.py +++ b/test/registered/disaggregation/test_disaggregation_basic.py @@ -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}", diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index 6ecf1d3eb..6cddeb03b 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -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), diff --git a/test/registered/unit/mem_cache/test_decode_retraction_backup.py b/test/registered/unit/mem_cache/test_decode_retraction_backup.py new file mode 100644 index 000000000..052ded614 --- /dev/null +++ b/test/registered/unit/mem_cache/test_decode_retraction_backup.py @@ -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() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 84d472924..5d462fe83 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -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):