[PD] Preserve decode KV across retraction in HiCache (#34801)
Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user