[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:
co-authored by
Zhiqiang Xie
parent
5d6b35eabb
commit
bcf298c28c
@@ -106,6 +106,8 @@ class RequestFuncOutput:
|
|||||||
error: str = ""
|
error: str = ""
|
||||||
output_len: int = 0
|
output_len: int = 0
|
||||||
start_time: float = 0.0
|
start_time: float = 0.0
|
||||||
|
cached_tokens: int = 0
|
||||||
|
cached_tokens_details: Optional[Dict[str, Any]] = None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def init_new(request_func_input: RequestFuncInput):
|
def init_new(request_func_input: RequestFuncInput):
|
||||||
@@ -229,6 +231,19 @@ async def async_request_trt_llm(
|
|||||||
return output
|
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
|
# set ignore_eos True by default
|
||||||
async def async_request_openai_completions(
|
async def async_request_openai_completions(
|
||||||
request_func_input: RequestFuncInput,
|
request_func_input: RequestFuncInput,
|
||||||
@@ -302,6 +317,9 @@ async def async_request_openai_completions(
|
|||||||
else:
|
else:
|
||||||
data = json.loads(chunk)
|
data = json.loads(chunk)
|
||||||
|
|
||||||
|
if getattr(args, "cache_report", False):
|
||||||
|
_extract_cache_from_sglext(data, output)
|
||||||
|
|
||||||
# NOTE: Some completion API might have a last
|
# NOTE: Some completion API might have a last
|
||||||
# usage summary response without a token so we
|
# usage summary response without a token so we
|
||||||
# want to check a token was generated
|
# 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(
|
output.output_len = response_json.get("usage", {}).get(
|
||||||
"completion_tokens", output_len
|
"completion_tokens", output_len
|
||||||
)
|
)
|
||||||
|
if getattr(args, "cache_report", False):
|
||||||
|
_extract_cache_from_sglext(response_json, output)
|
||||||
else:
|
else:
|
||||||
# Streaming response
|
# Streaming response
|
||||||
async for chunk_bytes in response.content:
|
async for chunk_bytes in response.content:
|
||||||
@@ -474,6 +494,9 @@ async def async_request_openai_chat_completions(
|
|||||||
"completion_tokens", output_len
|
"completion_tokens", output_len
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if getattr(args, "cache_report", False):
|
||||||
|
_extract_cache_from_sglext(data, output)
|
||||||
|
|
||||||
choices = data.get("choices") or []
|
choices = data.get("choices") or []
|
||||||
if not choices:
|
if not choices:
|
||||||
continue
|
continue
|
||||||
@@ -675,6 +698,13 @@ async def async_request_sglang_generate(
|
|||||||
# NOTE: Some completion API might have a last
|
# NOTE: Some completion API might have a last
|
||||||
# usage summary response without a token so we
|
# usage summary response without a token so we
|
||||||
# want to check a token was generated
|
# 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"]:
|
if "text" in data and data["text"]:
|
||||||
timestamp = time.perf_counter()
|
timestamp = time.perf_counter()
|
||||||
generated_text = data["text"]
|
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("P95 ITL (ms):", metrics.p95_itl_ms))
|
||||||
print("{:<40} {:<10.2f}".format("P99 ITL (ms):", metrics.p99_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))
|
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)
|
print("=" * 50)
|
||||||
|
|
||||||
resp = requests.get(base_url + "/server_info", headers=get_auth_headers())
|
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_output_tokens_per_s": metrics.max_output_tokens_per_s,
|
||||||
"max_concurrent_requests": metrics.max_concurrent_requests,
|
"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:
|
else:
|
||||||
print(f"Error running benchmark for request rate: {request_rate}")
|
print(f"Error running benchmark for request rate: {request_rate}")
|
||||||
print("-" * 30)
|
print("-" * 30)
|
||||||
@@ -1703,6 +1796,12 @@ async def benchmark(
|
|||||||
"errors": [output.error for output in outputs],
|
"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
|
# Append results to a JSONL file
|
||||||
with open(output_file_name, "a") as file:
|
with open(output_file_name, "a") as file:
|
||||||
if args.output_details:
|
if args.output_details:
|
||||||
@@ -1778,6 +1877,9 @@ def run_benchmark(args_: argparse.Namespace):
|
|||||||
if not hasattr(args, "served_model_name"):
|
if not hasattr(args, "served_model_name"):
|
||||||
args.served_model_name = None
|
args.served_model_name = None
|
||||||
|
|
||||||
|
if not hasattr(args, "cache_report"):
|
||||||
|
args.cache_report = False
|
||||||
|
|
||||||
if getattr(args, "print_requests", False):
|
if getattr(args, "print_requests", False):
|
||||||
assert args.backend == "sglang-oai-chat" # only support this now
|
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:
|
if args.extra_request_body:
|
||||||
extra_request_body = json.loads(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
|
# Inject bootstrap fields for fake decode benchmarking
|
||||||
if getattr(args, "fake_prefill", False):
|
if getattr(args, "fake_prefill", False):
|
||||||
extra_request_body["bootstrap_host"] = FAKE_BOOTSTRAP_HOST
|
extra_request_body["bootstrap_host"] = FAKE_BOOTSTRAP_HOST
|
||||||
@@ -2259,6 +2368,12 @@ if __name__ == "__main__":
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Return routed experts.",
|
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("--seed", type=int, default=1, help="The random seed.")
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--disable-ignore-eos",
|
"--disable-ignore-eos",
|
||||||
|
|||||||
Reference in New Issue
Block a user