[HiCache & Bench] add cache hit breakdown in bench_serving (#22053)

Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com>
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
Vladislav Nosivskoy
2026-06-17 16:02:30 -07:00
committed by GitHub
co-authored by Zhiqiang Xie
parent 5d6b35eabb
commit bcf298c28c
+115
View File
@@ -106,6 +106,8 @@ class RequestFuncOutput:
error: str = ""
output_len: int = 0
start_time: float = 0.0
cached_tokens: int = 0
cached_tokens_details: Optional[Dict[str, Any]] = None
@staticmethod
def init_new(request_func_input: RequestFuncInput):
@@ -229,6 +231,19 @@ async def async_request_trt_llm(
return output
def _extract_cache_from_sglext(data, output):
"""Extract cache hit details from sglext in OAI-compatible responses."""
sglext = data.get("sglext") or {}
details = sglext.get("cached_tokens_details")
if details:
output.cached_tokens = (
(details.get("device") or 0)
+ (details.get("host") or 0)
+ (details.get("storage") or 0)
)
output.cached_tokens_details = details
# set ignore_eos True by default
async def async_request_openai_completions(
request_func_input: RequestFuncInput,
@@ -302,6 +317,9 @@ async def async_request_openai_completions(
else:
data = json.loads(chunk)
if getattr(args, "cache_report", False):
_extract_cache_from_sglext(data, output)
# NOTE: Some completion API might have a last
# usage summary response without a token so we
# want to check a token was generated
@@ -455,6 +473,8 @@ async def async_request_openai_chat_completions(
output.output_len = response_json.get("usage", {}).get(
"completion_tokens", output_len
)
if getattr(args, "cache_report", False):
_extract_cache_from_sglext(response_json, output)
else:
# Streaming response
async for chunk_bytes in response.content:
@@ -474,6 +494,9 @@ async def async_request_openai_chat_completions(
"completion_tokens", output_len
)
if getattr(args, "cache_report", False):
_extract_cache_from_sglext(data, output)
choices = data.get("choices") or []
if not choices:
continue
@@ -675,6 +698,13 @@ async def async_request_sglang_generate(
# NOTE: Some completion API might have a last
# usage summary response without a token so we
# want to check a token was generated
if getattr(args, "cache_report", False):
_meta = data.get("meta_info") or {}
output.cached_tokens = _meta.get("cached_tokens", 0)
output.cached_tokens_details = _meta.get(
"cached_tokens_details"
)
if "text" in data and data["text"]:
timestamp = time.perf_counter()
generated_text = data["text"]
@@ -1608,6 +1638,58 @@ async def benchmark(
print("{:<40} {:<10.2f}".format("P95 ITL (ms):", metrics.p95_itl_ms))
print("{:<40} {:<10.2f}".format("P99 ITL (ms):", metrics.p99_itl_ms))
print("{:<40} {:<10.2f}".format("Max ITL (ms):", metrics.max_itl_ms))
if args.cache_report:
total_prompt_tokens = 0
total_cached = 0
total_device = total_host = total_storage = 0
storage_backend_name = None
has_details = False
for o in outputs:
if not o.success:
continue
total_prompt_tokens += o.prompt_len
total_cached += o.cached_tokens
if o.cached_tokens_details:
has_details = True
total_device += o.cached_tokens_details.get("device") or 0
total_host += o.cached_tokens_details.get("host") or 0
s = o.cached_tokens_details.get("storage") or 0
if s:
total_storage += s
storage_backend_name = o.cached_tokens_details.get(
"storage_backend"
)
hit_rate = (
total_cached / total_prompt_tokens * 100 if total_prompt_tokens > 0 else 0.0
)
print("{s:{c}^{n}}".format(s="Cache Hit Details", n=50, c="-"))
print("{:<40} {:<10}".format("Total prompt tokens:", total_prompt_tokens))
print("{:<40} {:<10}".format("Total cached tokens:", total_cached))
if has_details and total_cached > 0:
print("{:<40} {:<10}".format(" Device:", total_device))
print("{:<40} {:<10}".format(" Host:", total_host))
if total_storage > 0:
label = (
f" Storage ({storage_backend_name}):"
if storage_backend_name
else " Storage:"
)
print("{:<40} {:<10}".format(label, total_storage))
print("{:<40} {:.1f}%".format("Cache hit rate:", hit_rate))
if has_details and total_cached > 0:
device_pct = total_device / total_cached * 100
host_pct = total_host / total_cached * 100
print("{:<40} {:.1f}%".format(" Device:", device_pct))
print("{:<40} {:.1f}%".format(" Host:", host_pct))
if total_storage > 0:
storage_pct = total_storage / total_cached * 100
label = (
f" Storage ({storage_backend_name}):"
if storage_backend_name
else " Storage:"
)
print("{:<40} {:.1f}%".format(label, storage_pct))
print("=" * 50)
resp = requests.get(base_url + "/server_info", headers=get_auth_headers())
@@ -1672,6 +1754,17 @@ async def benchmark(
"max_output_tokens_per_s": metrics.max_output_tokens_per_s,
"max_concurrent_requests": metrics.max_concurrent_requests,
}
if args.cache_report:
result["cache_report"] = {
"total_prompt_tokens": total_prompt_tokens,
"total_cached_tokens": total_cached,
"cache_hit_rate_pct": round(hit_rate, 2),
"device_cached_tokens": total_device if has_details else None,
"host_cached_tokens": total_host if has_details else None,
"storage_cached_tokens": (total_storage if total_storage > 0 else None),
"storage_backend": storage_backend_name,
}
else:
print(f"Error running benchmark for request rate: {request_rate}")
print("-" * 30)
@@ -1703,6 +1796,12 @@ async def benchmark(
"errors": [output.error for output in outputs],
}
if args.cache_report:
result_details["cached_tokens"] = [o.cached_tokens for o in outputs]
result_details["cached_tokens_details"] = [
o.cached_tokens_details for o in outputs
]
# Append results to a JSONL file
with open(output_file_name, "a") as file:
if args.output_details:
@@ -1778,6 +1877,9 @@ def run_benchmark(args_: argparse.Namespace):
if not hasattr(args, "served_model_name"):
args.served_model_name = None
if not hasattr(args, "cache_report"):
args.cache_report = False
if getattr(args, "print_requests", False):
assert args.backend == "sglang-oai-chat" # only support this now
@@ -1792,6 +1894,13 @@ def run_benchmark(args_: argparse.Namespace):
if args.extra_request_body:
extra_request_body = json.loads(args.extra_request_body)
if args.cache_report:
sglang_backends = ("sglang", "sglang-native", "sglang-oai", "sglang-oai-chat")
if args.backend not in sglang_backends:
print("WARNING: --cache-report is only supported with sglang backends.")
elif args.backend in ("sglang-oai", "sglang-oai-chat"):
extra_request_body["return_cached_tokens_details"] = True
# Inject bootstrap fields for fake decode benchmarking
if getattr(args, "fake_prefill", False):
extra_request_body["bootstrap_host"] = FAKE_BOOTSTRAP_HOST
@@ -2259,6 +2368,12 @@ if __name__ == "__main__":
action="store_true",
help="Return routed experts.",
)
parser.add_argument(
"--cache-report",
action="store_true",
help="Collect and display cache hit statistics after the benchmark. "
"Supported with sglang backends (native, oai, oai-chat).",
)
parser.add_argument("--seed", type=int, default=1, help="The random seed.")
parser.add_argument(
"--disable-ignore-eos",