Add cache hit rate UT (#18566)

This commit is contained in:
Liangsheng Yin
2026-02-10 21:27:41 -08:00
committed by GitHub
parent 50f74285e9
commit cd90346a2b
5 changed files with 379 additions and 126 deletions
+9 -11
View File
@@ -36,20 +36,18 @@ class ContextWorkloadGenerator(WorkloadGenerator):
init_requests = [] init_requests = []
for i in range(num_requests): for i in range(num_requests):
context_id = self.dataset["queries"][i]["context"] context_id = self.dataset["queries"][i]["context"]
init_requests.append( # Tokenize the context + question to get input_ids
( prompt_text = (
i,
gen_payload(
self.dataset["contexts"][context_id] self.dataset["contexts"][context_id]
+ self.dataset["queries"][i]["question"], + self.dataset["queries"][i]["question"]
len(
self.tokenizer(
self.dataset["queries"][i]["reference_answer"]
)["input_ids"]
),
),
) )
input_ids = self.tokenizer.encode(prompt_text)
output_len = len(
self.tokenizer(self.dataset["queries"][i]["reference_answer"])[
"input_ids"
]
) )
init_requests.append((i, gen_payload(input_ids, output_len)))
self.ready_queue = ReadyQueue(init_requests=init_requests) self.ready_queue = ReadyQueue(init_requests=init_requests)
self.response_queue = queue.Queue() self.response_queue = queue.Queue()
+14 -106
View File
@@ -6,21 +6,13 @@ import random
import threading import threading
import time import time
from datetime import datetime from datetime import datetime
from typing import Optional
import aiohttp
import numpy as np import numpy as np
import requests import requests
from tqdm.asyncio import tqdm from tqdm.asyncio import tqdm
from sglang.bench_serving import ( from sglang.bench_serving import get_tokenizer, sample_random_requests
RequestFuncOutput, from sglang.test.kits.cache_hit_kit import async_request_sglang_generate, gen_payload
get_tokenizer,
remove_prefix,
sample_random_requests,
)
AIOHTTP_TIMEOUT = aiohttp.ClientTimeout(total=20 * 60 * 60)
def parse_args(): def parse_args():
@@ -143,95 +135,6 @@ def parse_args():
return parser.parse_args() return parser.parse_args()
async def async_request_sglang_generate(
payload,
url,
pbar: Optional[tqdm] = None,
):
"""
Sends a streaming request to the server. Gathers text token-by-token.
"""
async with aiohttp.ClientSession(timeout=AIOHTTP_TIMEOUT) as session:
headers = {}
generated_text = ""
ttft = 0.0
st = time.perf_counter()
most_recent_timestamp = st
output = RequestFuncOutput()
try:
async with session.post(url=url, json=payload, headers=headers) as response:
if response.status == 200:
prompt_tokens = 0
cached_tokens = 0
async for chunk_bytes in response.content:
chunk_bytes = chunk_bytes.strip()
if not chunk_bytes:
continue
chunk = remove_prefix(chunk_bytes.decode("utf-8"), "data: ")
latency = time.perf_counter() - st
if chunk == "[DONE]":
pass
else:
data = json.loads(chunk)
if data["text"]:
timestamp = time.perf_counter()
# First token
if ttft == 0.0:
ttft = time.perf_counter() - st
output.ttft = ttft
prompt_tokens = (data.get("meta_info") or {}).get(
"prompt_tokens", 0
)
cached_tokens = (data.get("meta_info") or {}).get(
"cached_tokens", 0
)
# Decoding phase
else:
output.itl.append(timestamp - most_recent_timestamp)
most_recent_timestamp = timestamp
generated_text = data["text"]
output.generated_text = generated_text
output.success = True
output.latency = latency
output.prompt_len = prompt_tokens
output.cached_tokens = cached_tokens
output.generated_len = len(output.itl) + 1
else:
output.error = response.reason or ""
output.success = False
except Exception as e:
output.success = False
output.error = str(e)
print(f"Request failed: {e}")
if pbar:
pbar.update(1)
return output
def gen_payload(prompt, output_len, lora_path=""):
payload = {
"text": prompt,
"sampling_params": {
"temperature": 0.0,
"max_new_tokens": output_len,
"ignore_eos": True,
},
"stream": True,
"stream_options": {"include_usage": True},
"lora_path": lora_path,
"return_logprob": False,
"logprob_start_len": -1,
}
return payload
def log_to_jsonl_file(data, file_path="performance_metrics.jsonl", tag=""): def log_to_jsonl_file(data, file_path="performance_metrics.jsonl", tag=""):
"""Append the data with a timestamp and tag to the specified JSONL file.""" """Append the data with a timestamp and tag to the specified JSONL file."""
timestamped_data = {"timestamp": datetime.now().isoformat(), "tag": tag, **data} timestamped_data = {"timestamp": datetime.now().isoformat(), "tag": tag, **data}
@@ -286,6 +189,7 @@ class WorkloadGenerator:
self.sent_requests = 0 self.sent_requests = 0
self.completed_requests = 0 self.completed_requests = 0
# Use return_text=False to get token ids instead of text
self.candidate_inputs = sample_random_requests( self.candidate_inputs = sample_random_requests(
input_len=args.request_length, input_len=args.request_length,
output_len=args.output_length, output_len=args.output_length,
@@ -294,8 +198,10 @@ class WorkloadGenerator:
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
dataset_path=args.dataset_path, dataset_path=args.dataset_path,
random_sample=not args.disable_random_sample, random_sample=not args.disable_random_sample,
return_text=False,
) )
self.candidate_inputs = [i.prompt for i in self.candidate_inputs] # r.prompt is now List[int] when return_text=False
self.candidate_inputs = [list(i.prompt) for i in self.candidate_inputs]
if args.sub_question_input_length != 0: if args.sub_question_input_length != 0:
sub_question_input_length = args.sub_question_input_length sub_question_input_length = args.sub_question_input_length
@@ -310,6 +216,7 @@ class WorkloadGenerator:
tokenizer=self.tokenizer, tokenizer=self.tokenizer,
dataset_path=args.dataset_path, dataset_path=args.dataset_path,
random_sample=not args.disable_random_sample, random_sample=not args.disable_random_sample,
return_text=False,
) )
init_requests = [ init_requests = [
@@ -321,8 +228,9 @@ class WorkloadGenerator:
) )
for i in range(args.num_clients) for i in range(args.num_clients)
] ]
# history now stores List[int] (token ids) for each client
self.client_records = { self.client_records = {
i: {"round": 0, "history": init_requests[i][1]["text"]} i: {"round": 0, "history": list(self.candidate_inputs[i])}
for i in range(args.num_clients) for i in range(args.num_clients)
} }
self.ready_queue = ReadyQueue( self.ready_queue = ReadyQueue(
@@ -408,7 +316,8 @@ class WorkloadGenerator:
) # Block until response is available ) # Block until response is available
if not response.success: if not response.success:
raise ValueError(f"Request failed with error: {response.error}") raise ValueError(f"Request failed with error: {response.error}")
self.client_records[client_id]["history"] += response.generated_text # Use output_ids (token ids) instead of generated_text
self.client_records[client_id]["history"].extend(response.output_ids)
current_round = self.client_records[client_id]["round"] current_round = self.client_records[client_id]["round"]
self.client_records[client_id]["round"] += 1 self.client_records[client_id]["round"] += 1
self.performance_metrics["ttft"].append(response.ttft) self.performance_metrics["ttft"].append(response.ttft)
@@ -435,10 +344,9 @@ class WorkloadGenerator:
self.completed_requests += 1 self.completed_requests += 1
if self.client_records[client_id]["round"] < self.num_rounds: if self.client_records[client_id]["round"] < self.num_rounds:
# append new request to client's history # Append sub-question token ids to client's history
self.client_records[client_id][ sub_q_ids = list(self.sub_question_inputs.pop().prompt)
"history" self.client_records[client_id]["history"].extend(sub_q_ids)
] += self.sub_question_inputs.pop().prompt
new_req = ( new_req = (
client_id, client_id,
gen_payload( gen_payload(
+8 -7
View File
@@ -1505,12 +1505,12 @@ def sample_custom_requests(
return filtered_dataset return filtered_dataset
def compute_random_lens(full_len: int, range_ratio: float, num: int): def compute_random_lens(full_len: int, range_ratio: float, num: int) -> List[int]:
return np.random.randint( return np.random.randint(
max(int(full_len * range_ratio), 1), max(int(full_len * range_ratio), 1),
full_len + 1, full_len + 1,
size=num, size=num,
) ).tolist()
def sample_random_requests( def sample_random_requests(
@@ -1597,8 +1597,8 @@ def sample_random_requests(
input_requests.append( input_requests.append(
DatasetRow( DatasetRow(
prompt=input_content, prompt=input_content,
prompt_len=int(input_lens[i]), prompt_len=input_lens[i],
output_len=int(output_lens[i]), output_len=output_lens[i],
) )
) )
else: else:
@@ -1606,8 +1606,9 @@ def sample_random_requests(
offsets = np.random.randint(0, tokenizer.vocab_size, size=num_prompts) offsets = np.random.randint(0, tokenizer.vocab_size, size=num_prompts)
input_requests = [] input_requests = []
for i in range(num_prompts): for i in range(num_prompts):
# Use int() to convert numpy.int64 to native Python int for JSON serialization
input_content = [ input_content = [
(offsets[i] + i + j) % tokenizer.vocab_size int((offsets[i] + i + j) % tokenizer.vocab_size)
for j in range(input_lens[i]) for j in range(input_lens[i])
] ]
if return_text: if return_text:
@@ -1615,8 +1616,8 @@ def sample_random_requests(
input_requests.append( input_requests.append(
DatasetRow( DatasetRow(
prompt=input_content, prompt=input_content,
prompt_len=int(input_lens[i]), prompt_len=input_lens[i],
output_len=int(output_lens[i]), output_len=output_lens[i],
) )
) )
+300
View File
@@ -0,0 +1,300 @@
import asyncio
import json
import time
import aiohttp
import requests
from sglang.bench_serving import (
RequestFuncOutput,
get_tokenizer,
remove_prefix,
sample_random_requests,
)
AIOHTTP_TIMEOUT = aiohttp.ClientTimeout(total=20 * 60 * 60)
async def async_request_sglang_generate(
payload,
url,
pbar=None,
):
"""Send a streaming request to the server and collect cache metrics.
Returns a RequestFuncOutput with additional cached_tokens and output_ids attributes.
"""
async with aiohttp.ClientSession(timeout=AIOHTTP_TIMEOUT) as session:
headers = {}
generated_text = ""
all_output_ids = []
ttft = 0.0
st = time.perf_counter()
most_recent_timestamp = st
output = RequestFuncOutput()
try:
async with session.post(url=url, json=payload, headers=headers) as response:
if response.status == 200:
prompt_tokens = 0
cached_tokens = 0
async for chunk_bytes in response.content:
chunk_bytes = chunk_bytes.strip()
if not chunk_bytes:
continue
chunk = remove_prefix(chunk_bytes.decode("utf-8"), "data: ")
latency = time.perf_counter() - st
if chunk == "[DONE]":
pass
else:
data = json.loads(chunk)
# output_ids and text are always returned together
if data.get("output_ids"):
all_output_ids = data["output_ids"]
generated_text = data.get("text", "")
timestamp = time.perf_counter()
if ttft == 0.0:
ttft = time.perf_counter() - st
output.ttft = ttft
prompt_tokens = (data.get("meta_info") or {}).get(
"prompt_tokens", 0
)
cached_tokens = (data.get("meta_info") or {}).get(
"cached_tokens", 0
)
else:
output.itl.append(timestamp - most_recent_timestamp)
most_recent_timestamp = timestamp
output.generated_text = generated_text
output.output_ids = all_output_ids
output.success = True
output.latency = latency
output.prompt_len = prompt_tokens
output.cached_tokens = cached_tokens
output.generated_len = len(output.itl) + 1
else:
output.error = response.reason or ""
output.success = False
except Exception as e:
output.success = False
output.error = str(e)
print(f"Request failed: {e}")
if pbar:
pbar.update(1)
return output
def gen_payload(input_ids, output_len, lora_path=""):
return {
"input_ids": input_ids,
"sampling_params": {
"temperature": 0.0,
"max_new_tokens": output_len,
"ignore_eos": True,
},
"stream": True,
"stream_options": {"include_usage": True},
"lora_path": lora_path,
"return_logprob": False,
"logprob_start_len": -1,
}
async def _send_round(
payloads,
url,
max_parallel,
):
"""Send a batch of payloads concurrently with concurrency limit."""
semaphore = asyncio.Semaphore(max_parallel)
async def _send_one(payload):
async with semaphore:
return await async_request_sglang_generate(payload, url)
tasks = [asyncio.create_task(_send_one(p)) for p in payloads]
return await asyncio.gather(*tasks)
def _get_page_size(base_url: str) -> int:
"""Query server for page_size used by radix cache."""
try:
resp = requests.get(f"{base_url}/get_server_info", timeout=10)
resp.raise_for_status()
info = resp.json()
return info.get("page_size", 1)
except Exception:
return 1
def run_multiturn_cache_hit_test(
base_url: str,
model_path: str,
num_clients: int = 8,
num_rounds: int = 3,
request_length: int = 256,
output_length: int = 32,
miss_tolerance: int = 1,
sub_question_input_length: int = 0,
lora_path: str = "",
dataset_path: str = "",
max_parallel: int = 64,
seed: int = 1,
) -> dict:
"""Run a multi-turn workload and verify cache hit rate.
Sends requests in round-barrier mode: all clients complete round i
before round i+1 starts, ensuring deterministic cache state.
The expected cache hit rate is self-computed from the workload structure:
- Round 0: expected cached_tokens = 0 (cold start after flush)
- Round r (r >= 1): each client's prefix from round r-1 should be cached,
minus up to previous round's (prompt_len + decoding output - miss_tolerance) // page * page.
Returns metrics dict with per-round and overall cache_hit_rate.
"""
import random
import numpy as np
random.seed(seed)
np.random.seed(seed)
generate_url = f"{base_url}/generate"
page_size = _get_page_size(base_url)
# Flush cache for clean state
requests.post(f"{base_url}/flush_cache")
time.sleep(1)
# Resolve sub-question length (0 means same as request_length)
effective_sub_len = (
sub_question_input_length if sub_question_input_length != 0 else request_length
)
# Sample initial prompts and sub-question prompts as token ids
tokenizer = get_tokenizer(model_path)
initial_inputs = sample_random_requests(
input_len=request_length,
output_len=output_length,
num_prompts=num_clients,
range_ratio=1.0,
tokenizer=tokenizer,
dataset_path=dataset_path,
return_text=False,
)
# r.prompt is now List[int] when return_text=False
initial_token_ids = [list(r.prompt) for r in initial_inputs]
sub_question_inputs = sample_random_requests(
input_len=effective_sub_len,
output_len=output_length,
num_prompts=num_clients * max(num_rounds - 1, 1),
range_ratio=1.0,
tokenizer=tokenizer,
dataset_path=dataset_path,
return_text=False,
)
sub_question_token_ids = [list(r.prompt) for r in sub_question_inputs]
# Per-round metrics and per-client tracking for expected cache computation
round_metrics = {
i: {"prompt_len": [], "cached_tokens": [], "ttft": []}
for i in range(num_rounds)
}
# Track the previous round's prompt_len per client to compute expected cache
prev_prompt_lens = [0] * num_clients
# histories now stores List[int] (token ids) for each client
histories = [list(ids) for ids in initial_token_ids]
sub_idx = 0
for round_num in range(num_rounds):
payloads = [gen_payload(h, output_length, lora_path) for h in histories]
responses = asyncio.run(_send_round(payloads, generate_url, max_parallel))
for i, resp in enumerate(responses):
assert resp.success, f"Round {round_num}, client {i} failed: {resp.error}"
round_metrics[round_num]["prompt_len"].append(resp.prompt_len)
round_metrics[round_num]["cached_tokens"].append(resp.cached_tokens)
round_metrics[round_num]["ttft"].append(resp.ttft)
# Verify cache hit against expected value
if round_num == 0:
# Cold start: no cache expected
expected_cached = 0
else:
# Previous round's prompt + output are in cache.
# Radix cache aligns to page_size, so the last partial page
# may not be cached.
cacheable = prev_prompt_lens[i] + output_length - miss_tolerance
expected_cached = (cacheable // page_size) * page_size
msg = (
f"Round {round_num}, client {i}: "
f"cached_tokens={resp.cached_tokens}, "
f"expected>={expected_cached} "
f"(prev_prompt={prev_prompt_lens[i]}, "
f"output={output_length}, page_size={page_size})"
)
print(msg)
assert resp.cached_tokens >= expected_cached
# Record this round's prompt_len for next round's expected calc
prev_prompt_lens[i] = resp.prompt_len
# Accumulate history for next round using output_ids (token ids)
histories[i].extend(resp.output_ids)
if round_num < num_rounds - 1:
histories[i].extend(sub_question_token_ids[sub_idx])
sub_idx += 1
# Compute per-round and overall cache hit rate
total_prompt = 0
total_cached = 0
result = {"rounds": {}, "overall": {}}
for r in range(num_rounds):
rm = round_metrics[r]
r_prompt = sum(rm["prompt_len"])
r_cached = sum(rm["cached_tokens"])
r_hit_rate = r_cached / r_prompt if r_prompt > 0 else 0.0
r_avg_ttft = sum(rm["ttft"]) / len(rm["ttft"]) if rm["ttft"] else 0.0
result["rounds"][f"round_{r}"] = {
"cache_hit_rate": r_hit_rate,
"average_ttft": r_avg_ttft,
"total_prompt_tokens": r_prompt,
"total_cached_tokens": r_cached,
"request_count": len(rm["ttft"]),
}
total_prompt += r_prompt
total_cached += r_cached
print(
f" Round {r}: cache_hit_rate={r_hit_rate:.4f}, "
f"avg_ttft={r_avg_ttft:.4f}s, "
f"cached={r_cached}/{r_prompt} tokens"
)
overall_hit_rate = total_cached / total_prompt if total_prompt > 0 else 0.0
result["overall"] = {
"cache_hit_rate": overall_hit_rate,
"total_prompt_tokens": total_prompt,
"total_cached_tokens": total_cached,
}
print(f" Overall cache_hit_rate={overall_hit_rate:.4f}")
return result
@@ -0,0 +1,46 @@
import unittest
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.cache_hit_kit import run_multiturn_cache_hit_test
from sglang.test.test_utils import (
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=120, suite="stage-b-test-small-1-gpu")
MODEL = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
class TestRadixCacheHit(CustomTestCase):
@classmethod
def setUpClass(cls):
cls.model = MODEL
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
cls.model,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
)
@classmethod
def tearDownClass(cls):
kill_process_tree(cls.process.pid)
def test_multiturn_cache_hit(self):
run_multiturn_cache_hit_test(
base_url=self.base_url,
model_path=self.model,
num_clients=8,
num_rounds=6,
request_length=289,
output_length=367,
)
if __name__ == "__main__":
unittest.main()