[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
@@ -225,13 +225,15 @@ class TestDisaggregationMooncakeFailure(PDDisaggregationServerBase):
class TestDisaggregationMooncakeSpec(
JSONConstrainedMixin, SpecGrammarKit, PDDisaggregationServerBase
):
min_retraction_accept_length = 1.3
@classmethod
def setUpClass(cls):
super().setUpClass()
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
spec_args = [
"--speculative-algorithm",
"EAGLE",
"EAGLE3",
"--speculative-draft-model-path",
DEFAULT_DRAFT_MODEL_EAGLE3,
"--speculative-num-steps",
@@ -245,9 +247,51 @@ class TestDisaggregationMooncakeSpec(
"--dtype=float16",
]
cls.extra_prefill_args = spec_args
cls.extra_decode_args = spec_args
cls.extra_decode_args = [
*spec_args,
"--disaggregation-decode-retraction-backup",
"host_pool",
]
cls.extra_decode_env = {"SGLANG_TEST_RETRACT": "true"}
cls.launch_all()
def test_host_pool_retraction_preserves_spec_acceptance(self):
prompts = [
f"Request {i}: explain how speculative decoding works. " * 4
for i in range(4)
]
response = requests.post(
self.lb_url + "/generate",
json={
"text": prompts,
"sampling_params": {
"temperature": 0,
"ignore_eos": True,
"max_new_tokens": 64,
},
},
)
response.raise_for_status()
results = response.json()
retracted_results = [
result for result in results if result["meta_info"]["num_retractions"] > 0
]
retraction_count = sum(
result["meta_info"]["num_retractions"] for result in retracted_results
)
self.assertGreater(retraction_count, 0)
completion_tokens = sum(
result["meta_info"]["completion_tokens"] for result in retracted_results
)
verify_count = sum(
result["meta_info"]["spec_verify_ct"] for result in retracted_results
)
self.assertGreater(verify_count, 0)
accept_length = completion_tokens / verify_count
print(f"Retraction speculative {accept_length=:.4f}")
self.assertGreater(accept_length, self.min_retraction_accept_length)
def test_gsm8k(self):
args = SimpleNamespace(
base_url=f"http://{self.base_host}:{self.lb_port}",
@@ -12,6 +12,7 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.managers.schedule_batch import FINISH_ABORT
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -31,6 +32,14 @@ class FakeReceiver:
class TestDecodeQueueCleanup(CustomTestCase):
def test_paged_swa_retraction_resume_uses_physical_page_budget(self):
# resume_retracted_reqs reads the retraction backend off the disagg
# bag, so the case publishes a config instead of injecting one.
override = get_context().override_server_args(
disaggregation_decode_retraction_backup="cpu_tensor"
)
override.install()
self.addCleanup(override.restore)
page_size = 128
fill_len = 574
physical_tokens_per_req = 5 * page_size
@@ -42,6 +51,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
origin_input_ids=[0] * fill_len,
output_ids=[],
is_retracted=True,
retraction_backup=None,
load_kv_cache=MagicMock(),
)
for i in range(4)
@@ -52,6 +62,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue.num_reserved_decode_tokens = 0
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
queue.tree_cache = MagicMock()
queue.scheduler = SimpleNamespace(
sliding_window_size=2047,
server_args=SimpleNamespace(disable_radix_cache=True),
@@ -0,0 +1,173 @@
import unittest
from types import SimpleNamespace
import torch
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.mem_cache.hicache_storage import PoolName
from sglang.srt.mem_cache.kv_cache_builder import maybe_register_hicache_draft
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, ReqToTokenPool
from sglang.srt.mem_cache.unified_cache.components import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.speculative.base_spec_worker import (
HiCacheDraftMode,
HiCacheDraftPlan,
)
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
class TestDecodeRetractionBackup(unittest.TestCase):
pool_size = 32
num_tokens = 8
dtype = torch.bfloat16
device = "cuda"
def _make_pool(self, layer_num: int) -> MHATokenToKVPool:
return MHATokenToKVPool(
size=self.pool_size,
page_size=1,
head_num=2,
head_dim=64,
dtype=self.dtype,
layer_num=layer_num,
device=self.device,
enable_memory_saver=False,
)
def _seed_pool(
self, pool: MHATokenToKVPool, indices: torch.Tensor, base: int
) -> None:
for layer_id, (key, value) in enumerate(
zip(pool.k_buffer, pool.v_buffer, strict=True)
):
pattern = torch.arange(
key[indices].numel(), device=self.device, dtype=torch.float32
).reshape_as(key[indices])
key[indices] = (pattern + base + layer_id * 100).to(self.dtype)
value[indices] = (pattern + base + 50 + layer_id * 100).to(self.dtype)
@staticmethod
def _snapshot_pool(
pool: MHATokenToKVPool, indices: torch.Tensor
) -> list[tuple[torch.Tensor, torch.Tensor]]:
return [
(key[indices].clone(), value[indices].clone())
for key, value in zip(pool.k_buffer, pool.v_buffer, strict=True)
]
def _assert_pool_equal(
self,
pool: MHATokenToKVPool,
indices: torch.Tensor,
expected: list[tuple[torch.Tensor, torch.Tensor]],
) -> None:
for (key, value), (expected_key, expected_value) in zip(
zip(pool.k_buffer, pool.v_buffer, strict=True), expected, strict=True
):
self.assertTrue(torch.equal(key[indices], expected_key))
self.assertTrue(torch.equal(value[indices], expected_value))
def test_restores_target_and_draft_kv(self):
server_args = ServerArgs(
model_path="dummy",
page_size=1,
hicache_ratio=1.0,
hicache_io_backend="kernel",
hicache_mem_layout="page_first",
)
set_global_server_args_for_scheduler(server_args)
req_to_token_pool = ReqToTokenPool(
size=2,
max_context_len=self.pool_size,
device=self.device,
enable_memory_saver=False,
)
target_pool = self._make_pool(layer_num=2)
allocator = TokenToKVPoolAllocator(
size=self.pool_size,
dtype=self.dtype,
device=self.device,
kvcache=target_pool,
need_sort=False,
)
params = CacheInitParams(
disable=True,
req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=allocator,
page_size=1,
is_eagle=True,
tree_components=(ComponentType.FULL,),
)
cache = UnifiedRadixCache(params)
cache.init_hicache(server_args, params)
self.addCleanup(cache.release_host_resources)
draft_pool = self._make_pool(layer_num=1)
maybe_register_hicache_draft(
tree_cache=cache,
draft_plan=HiCacheDraftPlan(
mode=HiCacheDraftMode.SIDECAR,
device_pools=(draft_pool,),
),
server_args=server_args,
page_size=1,
)
self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map)
cache.validate_retraction_host_capacity()
req = SimpleNamespace(
rid="request", req_pool_idx=None, seqlen=self.num_tokens + 1
)
self.assertIsNotNone(req_to_token_pool.alloc([req]))
source_indices = allocator.alloc(self.num_tokens)
self.assertIsNotNone(source_indices)
req_to_token_pool.write(
(req.req_pool_idx, slice(0, self.num_tokens)), source_indices
)
self._seed_pool(target_pool, source_indices, base=1000)
self._seed_pool(draft_pool, source_indices, base=3000)
target_expected = self._snapshot_pool(target_pool, source_indices)
draft_expected = self._snapshot_pool(draft_pool, source_indices)
host_free_before = cache.host_pool_group.available_size()
backup = cache.retraction_backup(req)
self.assertEqual(
{transfer.name for transfer in backup.pool_transfers or []},
{PoolName.DRAFT},
)
self.assertLess(cache.host_pool_group.available_size(), host_free_before)
for buffer in (*target_pool.k_buffer, *target_pool.v_buffer):
buffer.fill_(-1)
for buffer in (*draft_pool.k_buffer, *draft_pool.v_buffer):
buffer.fill_(-2)
allocator.free(source_indices)
blocker_indices = allocator.alloc(self.num_tokens)
destination_indices = allocator.alloc(self.num_tokens)
self.assertIsNotNone(blocker_indices)
self.assertIsNotNone(destination_indices)
self.assertFalse(torch.equal(source_indices, destination_indices))
req_to_token_pool.write(
(req.req_pool_idx, slice(0, self.num_tokens)), destination_indices
)
cache.retraction_restore(req, backup)
self._assert_pool_equal(target_pool, destination_indices, target_expected)
self._assert_pool_equal(draft_pool, destination_indices, draft_expected)
self.assertEqual(cache.host_pool_group.available_size(), host_free_before)
allocator.free(blocker_indices)
allocator.free(destination_indices)
req_to_token_pool.free(req)
if __name__ == "__main__":
unittest.main()
@@ -1331,6 +1331,27 @@ class TestHiCacheArgs(unittest.TestCase):
self.assertEqual(args.hicache_mem_layout, "page_first")
self.assertIsNone(args.decode_attention_backend)
def test_decode_offload_rejects_host_pool_retraction(self):
args = self._make_args(
disaggregation_mode="decode",
disaggregation_decode_enable_offload_kvcache=True,
hicache_storage_backend="file",
disaggregation_decode_retraction_backup="host_pool",
)
with self.assertRaisesRegex(ValueError, "mutually exclusive"):
args._handle_cache_compatibility()
def test_decode_offload_allows_cpu_tensor_retraction(self):
args = self._make_args(
disaggregation_mode="decode",
disaggregation_decode_enable_offload_kvcache=True,
hicache_storage_backend="file",
disaggregation_decode_retraction_backup="cpu_tensor",
)
args._handle_cache_compatibility()
class TestNgramExternalSamArgs(CustomTestCase):
def _make_dummy_ngram_args(self, **overrides):