[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, kv_to_page_indices,
page_align_floor, page_align_floor,
release_kv_cache, release_kv_cache,
retraction_discard,
retraction_restore,
) )
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
from sglang.srt.mem_cache.memory_pool import ( from sglang.srt.mem_cache.memory_pool import (
@@ -693,6 +695,12 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
def release_memory_occupation(self): def release_memory_occupation(self):
self.queue.clear() 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() self.retracted_queue.clear()
if hasattr(self.kv_manager, "deregister_buffer_to_engine"): if hasattr(self.kv_manager, "deregister_buffer_to_engine"):
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: if uses_swa_tail_prealloc:
swa_allocatable_tokens -= swa_required swa_allocatable_tokens -= swa_required
# load from cpu, release the cpu copy retraction_restore(
req.load_kv_cache(self.req_to_token_pool, self.token_to_kv_pool_allocator) req,
self.tree_cache,
self.req_to_token_pool,
self.token_to_kv_pool_allocator,
get_disagg().disaggregation_decode_retraction_backup,
)
self.retracted_queue = [ self.retracted_queue = [
entry entry
+28 -6
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
from sglang.srt.dllm.config import DllmConfig from sglang.srt.dllm.config import DllmConfig
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import ( from sglang.srt.runtime_context import (
get_disagg,
get_exec, get_exec,
get_schedule, get_schedule,
get_serving, get_serving,
@@ -97,9 +98,11 @@ from sglang.srt.mem_cache.base_prefix_cache import (
zero_match_result, zero_match_result,
) )
from sglang.srt.mem_cache.common import ( from sglang.srt.mem_cache.common import (
RetractionBackup,
evict_from_tree_cache, evict_from_tree_cache,
free_swa_out_of_window_slots, free_swa_out_of_window_slots,
release_kv_cache, release_kv_cache,
retraction_backup,
) )
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
@@ -878,6 +881,7 @@ class Req(ReqDllmMixin):
# For req-level memory management # For req-level memory management
self.kv_committed_len = 0 self.kv_committed_len = 0
self.kv: Optional[ReqKvInfo] = None self.kv: Optional[ReqKvInfo] = None
self.retraction_backup: Optional[RetractionBackup] = None
# for cross-encoder model # for cross-encoder model
self.token_type_ids = token_type_ids self.token_type_ids = token_type_ids
@@ -1712,19 +1716,24 @@ class Req(ReqDllmMixin):
self.req_pool_idx, : self.seqlen - 1 self.req_pool_idx, : self.seqlen - 1
] ]
# Copies over both the kv cache and mamba state if available # Copies over both the kv cache and mamba state if available
self.kv_cache_cpu = token_to_kv_pool_allocator.get_cpu_copy( self.retraction_backup = RetractionBackup(
cpu_tensors=token_to_kv_pool_allocator.get_cpu_copy(
token_indices, mamba_indices=self.mamba_pool_idx token_indices, mamba_indices=self.mamba_pool_idx
) )
)
def load_kv_cache(self, req_to_token_pool, token_to_kv_pool_allocator): 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[ token_indices = req_to_token_pool.req_to_token[
self.req_pool_idx, : self.seqlen - 1 self.req_pool_idx, : self.seqlen - 1
] ]
# Loads both the kv cache and mamba state if exists # Loads both the kv cache and mamba state if exists
token_to_kv_pool_allocator.load_cpu_copy( 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: def build_rebootstrap_payload(self) -> dict:
"""Build the prefill ``/generate`` payload that asks the original prefill """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) # Callers that will recompute the KV instead (PD true-retraction rebootstrap)
# pass offload_kv=False to skip the wasteful device->host copy. # pass offload_kv=False to skip the wasteful device->host copy.
if server_args.disaggregation_mode == "decode" and offload_kv: 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 # TODO (csy): for preempted requests, we may want to insert into the tree
release_kv_cache(req, tree_cache, is_insert=False) release_kv_cache(req, tree_cache, is_insert=False)
# NOTE(lsyin): we should use the newly evictable memory instantly. # NOTE(lsyin): we should use the newly evictable memory instantly.
@@ -2834,7 +2849,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
status_code=HTTPStatus.INTERNAL_SERVER_ERROR, status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
) )
reqs_to_abort.append(last_req) 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( logger.warning(
"retract_decode: aborted last request %s due to OOM", last_req.rid "retract_decode: aborted last request %s due to OOM", last_req.rid
) )
@@ -2889,7 +2904,13 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
) )
return sorted_indices 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( release_req(
req=self.reqs[idx], req=self.reqs[idx],
remaing_req_count=remaing_req_count, remaing_req_count=remaing_req_count,
@@ -2898,6 +2919,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
tree_cache=self.tree_cache, tree_cache=self.tree_cache,
hisparse_coordinator=self.hisparse_coordinator, hisparse_coordinator=self.hisparse_coordinator,
offload_kv=offload_kv,
) )
def prepare_encoder_info_decode(self): def prepare_encoder_info_decode(self):
+14 -4
View File
@@ -266,7 +266,11 @@ from sglang.srt.managers.utils import (
validate_input_length, validate_input_length,
) )
from sglang.srt.mem_cache import kv_cache_builder 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_executor.forward_batch_info import PPProxyTensors
from sglang.srt.model_loader.utils import get_resolved_model_impl from sglang.srt.model_loader.utils import get_resolved_model_impl
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
@@ -965,6 +969,9 @@ class Scheduler(
def init_memory_pools(self): def init_memory_pools(self):
"""Allocate KV cache pools for target and draft workers.""" """Allocate KV cache pools for target and draft workers."""
self.init_target_memory_pool() 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: if self.draft_worker is not None:
pool, allocator = self.tp_worker.get_memory_pool() pool, allocator = self.tp_worker.get_memory_pool()
self.draft_worker.alloc_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=}") logger.debug(f"Abort transfer queue request. {decode_req.req.rid=}")
decode_req.kv_receiver.abort() 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: if self.disagg_decode_prealloc_queue.retracted_queue:
remaining_retracted = [] remaining_retracted = []
for decode_req in self.disagg_decode_prealloc_queue.retracted_queue: for decode_req in self.disagg_decode_prealloc_queue.retracted_queue:
if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid): if recv_req.abort_all or decode_req.rid.startswith(recv_req.rid):
assert hasattr(decode_req, "kv_cache_cpu") retraction_discard(
del decode_req.kv_cache_cpu decode_req,
self.tree_cache,
get_disagg().disaggregation_decode_retraction_backup,
)
self.ipc_channels.send_to_tokenizer.send_output( self.ipc_channels.send_to_tokenizer.send_output(
AbortReq(rid=decode_req.rid), decode_req AbortReq(rid=decode_req.rid), decode_req
) )
+63 -1
View File
@@ -1,7 +1,7 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
from typing import TYPE_CHECKING from typing import TYPE_CHECKING, Any, NamedTuple, Optional, cast
import numpy as np import numpy as np
import torch 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.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.allocator.swa import SWATokenToKVPoolAllocator
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.hicache_storage import PoolTransfer
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import get_serving, get_spec from sglang.srt.runtime_context import get_serving, get_spec
from sglang.srt.utils.common import ceil_align from sglang.srt.utils.common import ceil_align
@@ -19,6 +20,7 @@ from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator 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. # 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 MAMBA_STATE_PER_REQ_PREFIX_CACHE = 3
@@ -29,6 +31,12 @@ MAMBA_STATE_PER_REQ_NO_CACHE = 1
logger = logging.getLogger(__name__) 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: def kv_to_page_indices(kv_indices: torch.Tensor, page_size: int) -> np.ndarray:
return (kv_indices[::page_size] // page_size).cpu().numpy() 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)) 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): 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 # the two resources currently have the same lifecycle, thus simplify logic below
assert (req.req_pool_idx is None) == (req.kv is None) 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, StorageMetricsCollector,
resolve_collector_class, resolve_collector_class,
) )
from sglang.srt.runtime_context import get_memory
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.cache_init_params import CacheInitParams
@@ -86,7 +87,7 @@ class HiRadixCache(RadixCache):
if isinstance(self.kv_cache, MHATokenToKVPool): if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)( self.token_to_kv_pool_host = get_mha_host_pool_cls(self.kv_cache)(
self.kv_cache, self.kv_cache,
server_args.hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size, server_args.hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, server_args.hicache_mem_layout,
@@ -104,7 +105,7 @@ class HiRadixCache(RadixCache):
_parallel = get_parallel() _parallel = get_parallel()
self.token_to_kv_pool_host = MLATokenToKVPoolHost( self.token_to_kv_pool_host = MLATokenToKVPoolHost(
self.kv_cache, self.kv_cache,
server_args.hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size, server_args.hicache_size,
self.page_size, self.page_size,
server_args.hicache_mem_layout, 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.pool_host.mla import MLATokenToKVPoolHost
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType 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: if TYPE_CHECKING:
import torch import torch
@@ -99,7 +99,7 @@ def build_kv_host_pool(
kwargs["dcp_rank"] = parallel.attn_dcp_rank kwargs["dcp_rank"] = parallel.attn_dcp_rank
return kv_host_pool_cls( return kv_host_pool_cls(
kv_pool, kv_pool,
server_args.hicache_ratio, get_memory().hicache_ratio,
server_args.hicache_size if host_size is None else host_size, server_args.hicache_size if host_size is None else host_size,
page_size, page_size,
server_args.hicache_mem_layout, server_args.hicache_mem_layout,
@@ -402,7 +402,7 @@ def _deepseek_v4_num_host_pages(
"DeepSeek V4 HiCache currently does not support --hicache-size; " "DeepSeek V4 HiCache currently does not support --hicache-size; "
"use --hicache-ratio instead." "use --hicache-ratio instead."
) )
ratio = server_args.hicache_ratio ratio = get_memory().hicache_ratio
full_host_pages = int(device_full_pages * ratio) full_host_pages = int(device_full_pages * ratio)
swa_host_pages = int(device_swa_pages * ratio) swa_host_pages = int(device_swa_pages * ratio)
return full_host_pages, swa_host_pages return full_host_pages, swa_host_pages
@@ -715,7 +715,7 @@ def build_hybrid_mamba_stack(
) )
mamba_host_pool = MambaPoolHost( mamba_host_pool = MambaPoolHost(
mamba_pool, mamba_pool,
server_args.hicache_ratio, get_memory().hicache_ratio,
mamba_host_size, mamba_host_size,
allocator_type=_get_allocator_type(server_args), allocator_type=_get_allocator_type(server_args),
layout=server_args.hicache_mem_layout, layout=server_args.hicache_mem_layout,
@@ -819,7 +819,7 @@ def build_hybrid_mamba_swa_stack(
) )
mamba_host_pool = MambaPoolHost( mamba_host_pool = MambaPoolHost(
mamba_pool, mamba_pool,
server_args.hicache_ratio, get_memory().hicache_ratio,
mamba_host_size, mamba_host_size,
allocator_type=server_args.hicache_storage_backend, allocator_type=server_args.hicache_storage_backend,
layout=server_args.hicache_mem_layout, 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.environ import envs
from sglang.srt.managers.mm_schedule import init_mm_embedding_cache 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.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.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.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: if TYPE_CHECKING:
@@ -130,6 +139,57 @@ def _register_legacy_hicache_draft(
tree_cache.cache_controller.set_draft_kv_pool(pool, draft_host_pool) 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( def build_kv_cache(
*, *,
server_args: ServerArgs, 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() 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 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 ( disable_radix_cache = server_args.disable_radix_cache or (
model_config.is_multimodal and uses_transformers_backend 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( maybe_register_hicache_draft(
tree_cache=tree_cache, tree_cache=tree_cache,
draft_plan=hicache_draft_plan, draft_plan=hicache_draft_plan,
@@ -268,6 +332,14 @@ def build_kv_cache(
page_size=page_size, 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() embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024) 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.environ import envs
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.cache_init_params import CacheInitParams 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 from sglang.srt.utils.tensor_bridge import use_mlx
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -82,6 +82,12 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
server_args = ctx.server_args server_args = ctx.server_args
params = ctx.params 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 ctx.effective_chunked_prefill_size is not None and ctx.disable_radix_cache:
if not ctx.is_hybrid_swa: if not ctx.is_hybrid_swa:
from sglang.srt.mem_cache.chunk_cache import ChunkCache from sglang.srt.mem_cache.chunk_cache import ChunkCache
@@ -167,6 +173,12 @@ def _create_unified_radix_cache(
params: CacheInitParams, params: CacheInitParams,
) -> BasePrefixCache: ) -> BasePrefixCache:
"""Initialize a UnifiedRadixCache with proper components and optional HiCache.""" """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_cache.components import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
@@ -197,7 +209,10 @@ def _create_unified_radix_cache(
ComponentType.MAMBA: MlxAuxiliaryStateComponent, ComponentType.MAMBA: MlxAuxiliaryStateComponent,
} }
cache = UnifiedRadixCache(params) 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) cache.init_hicache(server_args, params)
ctx.tp_worker.register_hicache_layer_transfer_counter( ctx.tp_worker.register_hicache_layer_transfer_counter(
cache.cache_controller.layer_done_counter cache.cache_controller.layer_done_counter
@@ -232,6 +247,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
"--enable-session-radix-cache)." "--enable-session-radix-cache)."
) )
hicache_attached = cache.cache_controller is not None
streaming_wrapped = False streaming_wrapped = False
if ( if (
ctx.server_args.enable_streaming_session ctx.server_args.enable_streaming_session
@@ -244,12 +260,12 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
logger.info( logger.info(
"Tree cache initialized: source=%s impl=%s hybrid_swa=%s hybrid_ssm=%s " "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, source,
type(cache).__name__, type(cache).__name__,
ctx.is_hybrid_swa, ctx.is_hybrid_swa,
ctx.is_hybrid_ssm, ctx.is_hybrid_ssm,
ctx.enable_hierarchical_cache, hicache_attached,
streaming_wrapped, streaming_wrapped,
) )
return cache return cache
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging import logging
import threading import threading
import time import time
from dataclasses import replace
from queue import Empty, Queue from queue import Empty, Queue
from typing import TYPE_CHECKING, Iterator, NamedTuple, Optional, Sequence, TypeVar 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.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.cache_controller import CacheOperation
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache, BasePrefixCache,
DecLockRefParams, DecLockRefParams,
@@ -23,6 +25,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams, MatchPrefixParams,
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.common import RetractionBackup
from sglang.srt.mem_cache.hicache_storage import ( from sglang.srt.mem_cache.hicache_storage import (
PoolHitPolicy, PoolHitPolicy,
PoolName, PoolName,
@@ -32,7 +35,9 @@ from sglang.srt.mem_cache.hicache_storage import (
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
HybridCacheController, HybridCacheController,
) )
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool
from sglang.srt.mem_cache.radix_cache import RadixKey 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 ( from sglang.srt.mem_cache.unified_cache.cache_action import (
BackupKV, BackupKV,
CacheAction, CacheAction,
@@ -70,6 +75,7 @@ from sglang.srt.observability.metrics_collector import (
resolve_collector_class, resolve_collector_class,
) )
from sglang.srt.session.streaming_session import StreamingSession from sglang.srt.session.streaming_session import StreamingSession
from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.cache_controller import HiCacheAck from sglang.srt.managers.cache_controller import HiCacheAck
@@ -425,6 +431,9 @@ class UnifiedRadixCache(BasePrefixCache):
assert not result.cache_actions assert not result.cache_actions
return result return result
def is_chunk_cache(self) -> bool:
return self.disable
def insert(self, params: InsertParams) -> InsertResult: def insert(self, params: InsertParams) -> InsertResult:
if self.disable: if self.disable:
return InsertResult(prefix_len=0) return InsertResult(prefix_len=0)
@@ -921,6 +930,240 @@ class UnifiedRadixCache(BasePrefixCache):
self._free_values(result.device_frees, result.host_frees) self._free_values(result.device_frees, result.host_frees)
return result.tracker.get(component_type, 0) 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 ---- # ---- HiCache: Backup / LoadBack ----
def _execute_and_commit_kv_backup( def _execute_and_commit_kv_backup(
+57 -3
View File
@@ -2653,10 +2653,10 @@ class ServerArgs:
False False
) )
hicache_ratio: A[ hicache_ratio: A[
float, Optional[float],
"The ratio of the size of host KV cache memory pool to the size of device pool.", "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"), NS("memory"),
] = 2.0 ] = None
hicache_size: A[ hicache_size: A[
int, int,
"The size of host KV cache memory pool in gigabytes, which will override the hicache_ratio if set.", "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).", "Enable async KV cache offloading on decode server (PD mode).",
NS("disagg"), NS("disagg"),
] = False ] = 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[ num_reserved_decode_tokens: A[
int, int,
"Number of decode tokens that will have memory reserved when adding new request to the running batch.", "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_moe_runner_backend_alias()
self._handle_return_hidden_states_mode() self._handle_return_hidden_states_mode()
self._handle_media_url_security() self._handle_media_url_security()
self._handle_hicache_ratio_default()
if self.model_path.lower() in ["none", "dummy"]: if self.model_path.lower() in ["none", "dummy"]:
return return
@@ -7335,6 +7349,10 @@ class ServerArgs:
"workloads and backends may be supported in a future change." "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): def _handle_hicache(self):
"""Normalize hicache-related knobs into a valid runtime configuration. """Normalize hicache-related knobs into a valid runtime configuration.
@@ -7346,6 +7364,10 @@ class ServerArgs:
if not ( if not (
self.enable_hierarchical_cache self.enable_hierarchical_cache
or self.disaggregation_decode_enable_offload_kvcache or self.disaggregation_decode_enable_offload_kvcache
or (
self.disaggregation_mode == "decode"
and self.disaggregation_decode_retraction_backup in (None, "host_pool")
)
): ):
return return
@@ -8084,6 +8106,32 @@ class ServerArgs:
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False) envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
def _handle_cache_compatibility(self): 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: if self.enable_hierarchical_cache and self.disable_radix_cache:
raise ValueError( raise ValueError(
"The arguments enable-hierarchical-cache and disable-radix-cache are mutually exclusive " "The arguments enable-hierarchical-cache and disable-radix-cache are mutually exclusive "
@@ -8099,6 +8147,12 @@ class ServerArgs:
raise ValueError( raise ValueError(
"The argument disaggregation-decode-enable-offload-kvcache is only supported when hicache-storage-backend is provided." "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 # Validate the effective ratio: model branches may declare a reset
# (e.g. Step3p forces 1.0 under hierarchical cache) that supersedes # (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_memory_usage,
merge_graph_time_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: if TYPE_CHECKING:
from sglang.srt.managers.io_struct import ( 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 = self.target_worker.model_runner
target_model_runner.mtp_draft_device_pools = () target_model_runner.mtp_draft_device_pools = ()
spec_algorithm = target_model_runner.spec_algorithm 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() return HiCacheDraftPlan()
draft_runners = self._draft_model_runners() draft_runners = self._draft_model_runners()
@@ -225,13 +225,15 @@ class TestDisaggregationMooncakeFailure(PDDisaggregationServerBase):
class TestDisaggregationMooncakeSpec( class TestDisaggregationMooncakeSpec(
JSONConstrainedMixin, SpecGrammarKit, PDDisaggregationServerBase JSONConstrainedMixin, SpecGrammarKit, PDDisaggregationServerBase
): ):
min_retraction_accept_length = 1.3
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
super().setUpClass() super().setUpClass()
cls.model = DEFAULT_TARGET_MODEL_EAGLE3 cls.model = DEFAULT_TARGET_MODEL_EAGLE3
spec_args = [ spec_args = [
"--speculative-algorithm", "--speculative-algorithm",
"EAGLE", "EAGLE3",
"--speculative-draft-model-path", "--speculative-draft-model-path",
DEFAULT_DRAFT_MODEL_EAGLE3, DEFAULT_DRAFT_MODEL_EAGLE3,
"--speculative-num-steps", "--speculative-num-steps",
@@ -245,9 +247,51 @@ class TestDisaggregationMooncakeSpec(
"--dtype=float16", "--dtype=float16",
] ]
cls.extra_prefill_args = spec_args 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() 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): def test_gsm8k(self):
args = SimpleNamespace( args = SimpleNamespace(
base_url=f"http://{self.base_host}:{self.lb_port}", 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.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler 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.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
@@ -31,6 +32,14 @@ class FakeReceiver:
class TestDecodeQueueCleanup(CustomTestCase): class TestDecodeQueueCleanup(CustomTestCase):
def test_paged_swa_retraction_resume_uses_physical_page_budget(self): 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 page_size = 128
fill_len = 574 fill_len = 574
physical_tokens_per_req = 5 * page_size physical_tokens_per_req = 5 * page_size
@@ -42,6 +51,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
origin_input_ids=[0] * fill_len, origin_input_ids=[0] * fill_len,
output_ids=[], output_ids=[],
is_retracted=True, is_retracted=True,
retraction_backup=None,
load_kv_cache=MagicMock(), load_kv_cache=MagicMock(),
) )
for i in range(4) for i in range(4)
@@ -52,6 +62,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue.num_reserved_decode_tokens = 0 queue.num_reserved_decode_tokens = 0
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs)) queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size) queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
queue.tree_cache = MagicMock()
queue.scheduler = SimpleNamespace( queue.scheduler = SimpleNamespace(
sliding_window_size=2047, sliding_window_size=2047,
server_args=SimpleNamespace(disable_radix_cache=True), 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.assertEqual(args.hicache_mem_layout, "page_first")
self.assertIsNone(args.decode_attention_backend) 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): class TestNgramExternalSamArgs(CustomTestCase):
def _make_dummy_ngram_args(self, **overrides): def _make_dummy_ngram_args(self, **overrides):