diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 3ff2a29f3..f1e7d5259 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1914,7 +1914,8 @@ def release_req( tree_cache: BasePrefixCache, hisparse_coordinator: Optional[HiSparseCoordinator], offload_kv: bool = True, -) -> None: +) -> bool: + """Returns False when the KV backup failed and the request cannot be resumed.""" if hisparse_coordinator is not None and not req.finished(): hisparse_coordinator.retract_req(req) @@ -1922,8 +1923,9 @@ def release_req( # restored later without recompute (see resume_retracted_reqs/load_kv_cache). # Callers that will recompute the KV instead (PD true-retraction rebootstrap) # pass offload_kv=False to skip the wasteful device->host copy. + backup_saved = True if server_args.disaggregation_mode == "decode" and offload_kv: - retraction_backup( + backup_saved = retraction_backup( req, tree_cache, req_to_token_pool, @@ -1937,6 +1939,7 @@ def release_req( evict_from_tree_cache(tree_cache, num_tokens) req.reset_for_retract() + return backup_saved def retract_all( @@ -2820,6 +2823,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): sorted_indices = self._get_decode_retraction_order(self.reqs, server_args) retracted_reqs = [] + reqs_to_abort: List[Req] = [] first_iter = True while first_iter or ( not self.check_decode_mem(selected_indices=sorted_indices) @@ -2831,11 +2835,23 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): first_iter = False idx = sorted_indices.pop() req = self.reqs[idx] - retracted_reqs.append(req) # release memory and don't insert into the tree because we need the space instantly - self.release_req(idx, len(sorted_indices), server_args) + if self.release_req(idx, len(sorted_indices), server_args): + retracted_reqs.append(req) + else: + # The retraction host pool could not hold the backup and the + # device KV is already freed, so the request cannot resume. + req.to_finish = FINISH_ABORT( + "Retraction host KV pool exhausted. Aborting the request.", + status_code=HTTPStatus.INTERNAL_SERVER_ERROR, + ) + reqs_to_abort.append(req) + logger.warning( + "retract_decode: aborted request %s, retraction host pool " + "exhausted", + req.rid, + ) - reqs_to_abort: List[Req] = [] if len(sorted_indices) <= 1 and not self.check_decode_mem( selected_indices=sorted_indices ): @@ -2910,8 +2926,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): remaing_req_count: int, server_args: ServerArgs, offload_kv: bool = True, - ): - release_req( + ) -> bool: + return release_req( req=self.reqs[idx], remaing_req_count=remaing_req_count, server_args=server_args, diff --git a/python/sglang/srt/mem_cache/common.py b/python/sglang/srt/mem_cache/common.py index 4bc61811b..1ece6720f 100644 --- a/python/sglang/srt/mem_cache/common.py +++ b/python/sglang/srt/mem_cache/common.py @@ -144,17 +144,20 @@ def retraction_backup( req_to_token_pool: ReqToTokenPool, token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator, backend: str, -) -> None: +) -> bool: + """Returns False when the host pool cannot hold the backup; the caller + aborts the request since its KV cannot be preserved.""" if backend == "cpu_tensor": req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator) - return + return True if backend != "host_pool": raise ValueError(f"Unknown retraction backup backend: {backend}") if req.seqlen <= 1: - return + return True unified_cache = cast("UnifiedRadixCache", tree_cache) req.retraction_backup = unified_cache.retraction_backup(req) + return req.retraction_backup is not None def retraction_restore( diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 71faa8129..eb00c1441 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -139,6 +139,12 @@ def _register_legacy_hicache_draft( tree_cache.cache_controller.set_draft_kv_pool(pool, draft_host_pool) +# Host slots a backup-only retraction pool gets, as a fraction of the device +# pool. Sized well under 1.0 because a retraction burst touches a fraction of +# the device tokens; overflow aborts the request rather than pre-reserving. +BACKUP_ONLY_HICACHE_RATIO = 0.2 + + def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str: """Resolve the retraction backend onto the config bags and return it. @@ -180,10 +186,14 @@ def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str: 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 + # Only a decode server reaches resolution with the ratio unset. A + # backup-only pool can be small: retractions that overflow it abort their + # request instead of crashing the scheduler. Sharing the pool with + # HiCache keeps the standard default. + if backend == "host_pool" and not memory.enable_hierarchical_cache: + fields["hicache_ratio"] = BACKUP_ONLY_HICACHE_RATIO + else: + fields["hicache_ratio"] = 2.0 source = "kv_cache_builder.decode_retraction" get_context().override(source, **fields) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 5477edf2c..498be1329 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -1042,24 +1042,6 @@ class UnifiedRadixCache(BasePrefixCache): "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 @@ -1138,7 +1120,8 @@ class UnifiedRadixCache(BasePrefixCache): return 0 return self.evict_host(num_tokens) - def retraction_backup(self, req: Req) -> RetractionBackup: + def retraction_backup(self, req: Req) -> Optional[RetractionBackup]: + """Back up device KV to the host pool; None when it cannot fit after reclaim.""" assert req.seqlen > 1 device_indices, extra_transfers = self._retraction_device_transfers(req) @@ -1147,11 +1130,7 @@ class UnifiedRadixCache(BasePrefixCache): 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()}." - ) + return None resolved = self.cache_controller._resolve_pool_transfers_allocation( extra_transfers or None, @@ -1161,10 +1140,7 @@ class UnifiedRadixCache(BasePrefixCache): ) 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]}." - ) + return None backup = RetractionBackup( host_indices=host_indices, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 679f56bf6..e7c8d4c05 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2701,7 +2701,7 @@ class ServerArgs: ] = "cache" hicache_ratio: A[ Optional[float], - "The ratio of the size of host KV cache memory pool to the size of device pool. Defaults to 2.0 in cache mode, 1.2 in buffer_only mode, or 1.0 for host-pool decode retraction.", + "The ratio of the size of host KV cache memory pool to the size of device pool. Defaults to 2.0 in cache mode, 1.2 in buffer_only mode, or 0.2 for backup-only host-pool decode retraction.", NS("memory"), ] = None hicache_size: A[ diff --git a/test/registered/disaggregation/test_disaggregation_basic.py b/test/registered/disaggregation/test_disaggregation_basic.py index 0e1c75230..2812ce216 100644 --- a/test/registered/disaggregation/test_disaggregation_basic.py +++ b/test/registered/disaggregation/test_disaggregation_basic.py @@ -5,6 +5,7 @@ import threading import time import unittest import uuid +from concurrent.futures import ThreadPoolExecutor from types import SimpleNamespace from typing import Any @@ -13,6 +14,7 @@ import openai import requests from transformers import AutoTokenizer +from sglang.srt.mem_cache.kv_cache_builder import BACKUP_ONLY_HICACHE_RATIO from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.json_constrained_kit import JSONConstrainedMixin from sglang.test.kits.pause_generation_kit import PauseResumeInPlaceMixin @@ -20,6 +22,7 @@ from sglang.test.kits.spec_server_kits import SpecGrammarKit from sglang.test.run_eval import run_eval from sglang.test.server_fixtures.disaggregation_fixture import ( PDDisaggregationServerBase, + assert_process_healthy, ) from sglang.test.test_utils import ( DEFAULT_DRAFT_MODEL_EAGLE3, @@ -292,6 +295,102 @@ class TestDisaggregationMooncakeSpec( print(f"Retraction speculative {accept_length=:.4f}") self.assertGreater(accept_length, self.min_retraction_accept_length) + def test_oversized_backup_aborts_only_its_own_request(self): + # Backup-only host_pool retraction sizes the host pool at a fraction of the + # device pool, so a long enough request cannot be backed up. Derive the + # length from the running server rather than pinning pool sizes, which + # would change what the other cases in this class exercise. + info = requests.get(self.decode_url + "/get_server_info", timeout=30).json() + device_tokens = info["max_total_num_tokens"] + host_slots = int(device_tokens * BACKUP_ONLY_HICACHE_RATIO) + # Over the host pool, but still inside both the device pool and the model + # context — a pool far larger than the context would reject the request + # before it ever reaches retraction. + oversized_len = min(int(device_tokens * 0.4), info["max_req_input_len"] - 1024) + self.assertGreater( + oversized_len, + host_slots, + f"no prompt length both overflows the {host_slots}-slot host pool and " + f"fits the {info['max_req_input_len']}-token context", + ) + + def oversized_request(seed): + # Sent on its own: a batched /generate fails as a whole once any member + # aborts, which would hide the concurrent traffic's own outcome. + # Generate long enough to still be decoding when a forced retraction + # lands — a short request finishes first and is never retracted. + return requests.post( + self.lb_url + "/generate", + json={ + "input_ids": [seed] * oversized_len, + "sampling_params": {"max_new_tokens": 512, "ignore_eos": True}, + }, + timeout=900, + ) + + def ordinary_request(seed): + # Must still be decoding when the oversized prefill lands: retraction + # keeps one request, so a batch that has drained to a single entry is + # skipped entirely and nothing is ever picked. + return requests.post( + self.lb_url + "/generate", + json={ + "input_ids": [seed] * 512, + "sampling_params": {"max_new_tokens": 4096, "ignore_eos": True}, + }, + timeout=900, + ) + + # Retraction picks the request with the fewest generated tokens; the prompt + # length only breaks ties. The oversized request has by far the longest + # prefill, so in a fixed batch it enters decode last, holds the fewest + # tokens, and is picked first — but only while nothing newer arrives, which + # is why the eval below runs after these rather than alongside them. + with ThreadPoolExecutor(max_workers=4) as pool: + oversized = pool.submit(oversized_request, 233) + ordinary = [pool.submit(ordinary_request, 300 + i) for i in range(3)] + response = oversized.result() + neighbours = [f.result() for f in ordinary] + + # A 200 here means the request was never retracted, not that the abort path + # is broken, so surface the retraction count to tell the two apart. + meta = ( + response.json().get("meta_info", {}) if response.status_code == 200 else {} + ) + self.assertEqual( + response.status_code, + 500, + ( + f"expected an aborted backup; got num_retractions=" + f"{meta.get('num_retractions')} completion_tokens=" + f"{meta.get('completion_tokens')}" + if meta + else response.text + ), + ) + self.assertIn("Retraction host KV pool exhausted", response.text) + for neighbour in neighbours: + self.assertEqual(neighbour.status_code, 200, neighbour.text) + + # The abort must leave the scheduler serving, and ordinary traffic must stay + # correct afterwards — a leaked host slot or a damaged neighbour shows up as + # a wrong answer rather than merely a 200. + assert_process_healthy(self, "decode", self.process_decode, self.decode_url) + metrics = run_eval( + SimpleNamespace( + base_url=f"http://{self.base_host}:{self.lb_port}", + eval_name="gsm8k", + api="completion", + max_tokens=512, + num_examples=64, + num_threads=32, + ) + ) + print(f"Post-abort gsm8k metrics: {metrics}") + # Looser than the 200-example test_gsm8k bar above: 64 examples is a + # health check on the post-abort server, not an accuracy measurement. + self.assertGreater(metrics["score"], 0.62) + def test_gsm8k(self): args = SimpleNamespace( base_url=f"http://{self.base_host}:{self.lb_port}", diff --git a/test/registered/unit/mem_cache/test_decode_retraction_backup.py b/test/registered/unit/mem_cache/test_decode_retraction_backup.py index 052ded614..43b68d45b 100644 --- a/test/registered/unit/mem_cache/test_decode_retraction_backup.py +++ b/test/registered/unit/mem_cache/test_decode_retraction_backup.py @@ -5,6 +5,7 @@ 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.common import retraction_backup 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 @@ -71,11 +72,12 @@ class TestDecodeRetractionBackup(unittest.TestCase): self.assertTrue(torch.equal(key[indices], expected_key)) self.assertTrue(torch.equal(value[indices], expected_value)) - def test_restores_target_and_draft_kv(self): + def _build_cache(self, hicache_ratio: float): + """Bring up a UnifiedRadixCache with a draft sidecar over fresh pools.""" server_args = ServerArgs( model_path="dummy", page_size=1, - hicache_ratio=1.0, + hicache_ratio=hicache_ratio, hicache_io_backend="kernel", hicache_mem_layout="page_first", ) @@ -119,16 +121,61 @@ class TestDecodeRetractionBackup(unittest.TestCase): ) self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map) cache.validate_retraction_host_capacity() + return SimpleNamespace( + server_args=server_args, + req_to_token_pool=req_to_token_pool, + allocator=allocator, + target_pool=target_pool, + draft_pool=draft_pool, + cache=cache, + ) - 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) + def _admit_req(self, env, num_tokens: int): + req = SimpleNamespace(rid="request", req_pool_idx=None, seqlen=num_tokens + 1) + self.assertIsNotNone(env.req_to_token_pool.alloc([req])) + source_indices = env.allocator.alloc(num_tokens) self.assertIsNotNone(source_indices) - req_to_token_pool.write( - (req.req_pool_idx, slice(0, self.num_tokens)), source_indices + env.req_to_token_pool.write( + (req.req_pool_idx, slice(0, num_tokens)), source_indices ) + return req, source_indices + + def test_backup_declined_when_host_pool_too_small(self): + # A backup-only host pool is deliberately smaller than the device pool, + # so a large enough request cannot be preserved. + env = self._build_cache(hicache_ratio=0.1) + self.assertLess(env.cache.host_pool_group.available_size(), self.num_tokens) + + req, source_indices = self._admit_req(env, self.num_tokens) + host_free_before = env.cache.host_pool_group.available_size() + + self.assertIsNone(env.cache.retraction_backup(req)) + # The declined backup must not leak host slots. + self.assertEqual(env.cache.host_pool_group.available_size(), host_free_before) + + # This is the signal release_req propagates so retract_decode aborts. + self.assertFalse( + retraction_backup( + req, + env.cache, + env.req_to_token_pool, + env.allocator, + "host_pool", + ) + ) + + env.allocator.free(source_indices) + env.req_to_token_pool.free(req) + + def test_restores_target_and_draft_kv(self): + env = self._build_cache(hicache_ratio=1.0) + req_to_token_pool = env.req_to_token_pool + allocator = env.allocator + target_pool = env.target_pool + draft_pool = env.draft_pool + cache = env.cache + + req, source_indices = self._admit_req(env, self.num_tokens) self._seed_pool(target_pool, source_indices, base=1000) self._seed_pool(draft_pool, source_indices, base=3000)