Add cache hit rate UT (#18566)
This commit is contained in:
@@ -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()
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user