Fix Mooncake serving benchmark trace rows (#37201)

This commit is contained in:
Shuwen Wang
2026-08-31 11:43:59 -07:00
committed by GitHub
parent 549166b819
commit 9e9d26a4af
+12 -4
View File
@@ -1340,7 +1340,7 @@ async def benchmark(
base_url: str,
model_id: str,
tokenizer: PreTrainedTokenizerBase,
input_requests: List[DatasetRow],
input_requests: List[Union[DatasetRow, Dict[str, Any]]],
request_rate: float,
max_concurrency: Optional[int],
disable_tqdm: bool,
@@ -1364,8 +1364,14 @@ async def benchmark(
else:
raise ValueError(f"Unknown backend: {backend}")
is_mooncake = args.dataset_name == "mooncake"
# Multi-turn iff prompt[0] is a valid per-round payload. Single-shot
# OpenAI messages (List[Dict]) is excluded since its first element is a dict.
if is_mooncake:
# Mooncake dataset rows are raw trace dictionaries. They are converted
# into DatasetRow objects by get_mooncake_request_over_time below.
is_multi_turn = False
else:
first_prompt = input_requests[0].prompt
is_multi_turn = (
isinstance(first_prompt, list)
@@ -1389,7 +1395,7 @@ async def benchmark(
print(f"Starting warmup with {warmup_requests} sequences...")
# Handle the data structure difference for the warmup request
if args.dataset_name == "mooncake":
if is_mooncake:
# For mooncake, input_requests is a list of dicts.
# We need to build a temporary DatasetRow for the warmup phase.
warmup_record = input_requests[0]
@@ -1493,7 +1499,7 @@ async def benchmark(
tasks: List[asyncio.Task] = []
pbar_total = len(input_requests)
if (
backend == "sglang" and args.dataset_name == "mooncake"
backend == "sglang" and is_mooncake
): # Assuming mooncake is mainly for sglang or similar backends
print("Using time-based Mooncake request scheduler, ignoring --request-rate.")
request_generator = get_mooncake_request_over_time(
@@ -1517,7 +1523,9 @@ async def benchmark(
lora_probs = None
pbar = None if disable_tqdm else tqdm(total=pbar_total)
benchmark_requests: List[DatasetRow] = []
async for request in request_generator:
benchmark_requests.append(request)
if lora_names is not None and len(lora_names) != 0:
if lora_request_distribution == "uniform":
lora_name = random.choice(lora_names)
@@ -1603,7 +1611,7 @@ async def benchmark(
# Compute metrics and print results
benchmark_duration = time.perf_counter() - benchmark_start_time
metrics, output_lens = calculate_metrics(
input_requests=None if is_multi_turn else input_requests,
input_requests=None if is_multi_turn else benchmark_requests,
outputs=outputs,
dur_s=benchmark_duration,
tokenizer=tokenizer,