[HiCache] Allow a retraction host pool smaller than the device pool (#35543)
Co-authored-by: cctry <cctry@fb.com>
This commit is contained in:
@@ -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}",
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user