[Bench] Add fixed-prompt mode and per-request spec accept length metrics (#30615)
This commit is contained in:
@@ -6,6 +6,7 @@ python3 -m sglang.test.run_eval --port 30000 --eval-name mmlu --num-examples 10
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import statistics
|
||||
import time
|
||||
|
||||
from sglang.test.simple_eval_common import (
|
||||
@@ -82,6 +83,7 @@ def run_eval_once(args, base_url: str, eval_obj: Eval) -> dict:
|
||||
**common_kwargs,
|
||||
reasoning_effort=getattr(args, "reasoning_effort", None),
|
||||
extra_body=extra_body if extra_body else None,
|
||||
record_meta_info=True,
|
||||
)
|
||||
|
||||
# Run eval
|
||||
@@ -92,6 +94,30 @@ def run_eval_once(args, base_url: str, eval_obj: Eval) -> dict:
|
||||
return result, latency, sampler
|
||||
|
||||
|
||||
def print_accept_length_summary(samplers: list) -> None:
|
||||
accept_lengths = [
|
||||
m["spec_accept_length"]
|
||||
for sampler in samplers
|
||||
for m in getattr(sampler, "_meta_infos", [])
|
||||
if m.get("spec_accept_length") is not None
|
||||
]
|
||||
print("=" * 20)
|
||||
if not accept_lengths:
|
||||
print(
|
||||
"Speculative decoding: no per-request spec_accept_length in responses "
|
||||
"(non-speculative server, or --api completion which lacks return_meta_info)."
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"Speculative accept length (per-request, from meta_info): "
|
||||
f"n={len(accept_lengths)} "
|
||||
f"mean={statistics.fmean(accept_lengths):.4f} "
|
||||
f"min={min(accept_lengths):.4f} "
|
||||
f"max={max(accept_lengths):.4f}"
|
||||
)
|
||||
print("=" * 20)
|
||||
|
||||
|
||||
def run_eval(args):
|
||||
# Lazy import to avoid circular dependency with test_utils
|
||||
from sglang.test.test_utils import dump_metric
|
||||
@@ -194,6 +220,7 @@ def run_eval(args):
|
||||
|
||||
if getattr(args, "repeat", 1) == 1:
|
||||
result, latency, sampler = run_eval_once(args, base_url, eval_obj)
|
||||
samplers = [sampler]
|
||||
metrics = result.metrics | {"score": result.score}
|
||||
metrics["latency"] = latency
|
||||
print(f"Total latency: {latency:.3f} s")
|
||||
@@ -229,9 +256,11 @@ def run_eval(args):
|
||||
scores_repeat = []
|
||||
latencies = []
|
||||
total_completion_tokens = 0
|
||||
samplers = []
|
||||
|
||||
for f in futures:
|
||||
result, latency, sampler = f.result()
|
||||
samplers.append(sampler)
|
||||
scores_repeat.append(result.score)
|
||||
latencies.append(latency)
|
||||
total_completion_tokens += sum(sampler._completion_tokens)
|
||||
@@ -266,6 +295,8 @@ def run_eval(args):
|
||||
|
||||
executor.shutdown()
|
||||
|
||||
print_accept_length_summary(samplers)
|
||||
|
||||
# Dump reports
|
||||
file_stem = f"{args.eval_name}_{sampler.model.replace('/', '_')}"
|
||||
report_filename = f"/tmp/{file_stem}.html"
|
||||
|
||||
@@ -95,6 +95,7 @@ class ChatCompletionSampler(SamplerBase):
|
||||
reasoning_effort: Optional[str] = None,
|
||||
max_tokens: int = 2048,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
record_meta_info: bool = False,
|
||||
):
|
||||
self.client = OpenAI(base_url=base_url, http_client=LargerHttpxClient())
|
||||
|
||||
@@ -110,8 +111,10 @@ class ChatCompletionSampler(SamplerBase):
|
||||
self.extra_body = extra_body
|
||||
self.image_format = "url"
|
||||
self._completion_tokens: list[int] = []
|
||||
self.record_meta_info = record_meta_info
|
||||
self._meta_infos: List[Dict[str, Any]] = []
|
||||
print(
|
||||
f"ChatCompletionSampler initialized with {self.system_message=} {self.temperature=} {self.max_tokens=} {self.reasoning_effort=} {self.extra_body=}"
|
||||
f"ChatCompletionSampler initialized with {self.system_message=} {self.temperature=} {self.max_tokens=} {self.reasoning_effort=} {self.extra_body=} {self.record_meta_info=}"
|
||||
)
|
||||
|
||||
def _handle_image(
|
||||
@@ -140,6 +143,9 @@ class ChatCompletionSampler(SamplerBase):
|
||||
message_list = [
|
||||
self._pack_message("system", self.system_message)
|
||||
] + message_list
|
||||
extra_body = self.extra_body
|
||||
if self.record_meta_info:
|
||||
extra_body = {**(self.extra_body or {}), "return_meta_info": True}
|
||||
trial = 0
|
||||
while trial < 6: # 126 seconds in total
|
||||
try:
|
||||
@@ -150,8 +156,12 @@ class ChatCompletionSampler(SamplerBase):
|
||||
top_p=self.top_p,
|
||||
max_tokens=self.max_tokens,
|
||||
reasoning_effort=self.reasoning_effort,
|
||||
extra_body=self.extra_body,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
if self.record_meta_info:
|
||||
meta_info = getattr(response.choices[0], "meta_info", None)
|
||||
if meta_info:
|
||||
self._meta_infos.append(meta_info)
|
||||
if response.usage and response.usage.completion_tokens is not None:
|
||||
self._completion_tokens.append(response.usage.completion_tokens)
|
||||
return response.choices[0].message.content or ""
|
||||
|
||||
Reference in New Issue
Block a user