Fix Mooncake serving benchmark trace rows (#37201)
This commit is contained in:
@@ -1340,7 +1340,7 @@ async def benchmark(
|
|||||||
base_url: str,
|
base_url: str,
|
||||||
model_id: str,
|
model_id: str,
|
||||||
tokenizer: PreTrainedTokenizerBase,
|
tokenizer: PreTrainedTokenizerBase,
|
||||||
input_requests: List[DatasetRow],
|
input_requests: List[Union[DatasetRow, Dict[str, Any]]],
|
||||||
request_rate: float,
|
request_rate: float,
|
||||||
max_concurrency: Optional[int],
|
max_concurrency: Optional[int],
|
||||||
disable_tqdm: bool,
|
disable_tqdm: bool,
|
||||||
@@ -1364,14 +1364,20 @@ async def benchmark(
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unknown backend: {backend}")
|
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
|
# 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.
|
# OpenAI messages (List[Dict]) is excluded since its first element is a dict.
|
||||||
first_prompt = input_requests[0].prompt
|
if is_mooncake:
|
||||||
is_multi_turn = (
|
# Mooncake dataset rows are raw trace dictionaries. They are converted
|
||||||
isinstance(first_prompt, list)
|
# into DatasetRow objects by get_mooncake_request_over_time below.
|
||||||
and bool(first_prompt)
|
is_multi_turn = False
|
||||||
and _normalize_round_messages(first_prompt[0]) is not None
|
else:
|
||||||
)
|
first_prompt = input_requests[0].prompt
|
||||||
|
is_multi_turn = (
|
||||||
|
isinstance(first_prompt, list)
|
||||||
|
and bool(first_prompt)
|
||||||
|
and _normalize_round_messages(first_prompt[0]) is not None
|
||||||
|
)
|
||||||
if is_multi_turn:
|
if is_multi_turn:
|
||||||
request_func = wrap_multi_turn_request_func(request_func, backend=backend)
|
request_func = wrap_multi_turn_request_func(request_func, backend=backend)
|
||||||
|
|
||||||
@@ -1389,7 +1395,7 @@ async def benchmark(
|
|||||||
print(f"Starting warmup with {warmup_requests} sequences...")
|
print(f"Starting warmup with {warmup_requests} sequences...")
|
||||||
|
|
||||||
# Handle the data structure difference for the warmup request
|
# 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.
|
# For mooncake, input_requests is a list of dicts.
|
||||||
# We need to build a temporary DatasetRow for the warmup phase.
|
# We need to build a temporary DatasetRow for the warmup phase.
|
||||||
warmup_record = input_requests[0]
|
warmup_record = input_requests[0]
|
||||||
@@ -1493,7 +1499,7 @@ async def benchmark(
|
|||||||
tasks: List[asyncio.Task] = []
|
tasks: List[asyncio.Task] = []
|
||||||
pbar_total = len(input_requests)
|
pbar_total = len(input_requests)
|
||||||
if (
|
if (
|
||||||
backend == "sglang" and args.dataset_name == "mooncake"
|
backend == "sglang" and is_mooncake
|
||||||
): # Assuming mooncake is mainly for sglang or similar backends
|
): # Assuming mooncake is mainly for sglang or similar backends
|
||||||
print("Using time-based Mooncake request scheduler, ignoring --request-rate.")
|
print("Using time-based Mooncake request scheduler, ignoring --request-rate.")
|
||||||
request_generator = get_mooncake_request_over_time(
|
request_generator = get_mooncake_request_over_time(
|
||||||
@@ -1517,7 +1523,9 @@ async def benchmark(
|
|||||||
lora_probs = None
|
lora_probs = None
|
||||||
|
|
||||||
pbar = None if disable_tqdm else tqdm(total=pbar_total)
|
pbar = None if disable_tqdm else tqdm(total=pbar_total)
|
||||||
|
benchmark_requests: List[DatasetRow] = []
|
||||||
async for request in request_generator:
|
async for request in request_generator:
|
||||||
|
benchmark_requests.append(request)
|
||||||
if lora_names is not None and len(lora_names) != 0:
|
if lora_names is not None and len(lora_names) != 0:
|
||||||
if lora_request_distribution == "uniform":
|
if lora_request_distribution == "uniform":
|
||||||
lora_name = random.choice(lora_names)
|
lora_name = random.choice(lora_names)
|
||||||
@@ -1603,7 +1611,7 @@ async def benchmark(
|
|||||||
# Compute metrics and print results
|
# Compute metrics and print results
|
||||||
benchmark_duration = time.perf_counter() - benchmark_start_time
|
benchmark_duration = time.perf_counter() - benchmark_start_time
|
||||||
metrics, output_lens = calculate_metrics(
|
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,
|
outputs=outputs,
|
||||||
dur_s=benchmark_duration,
|
dur_s=benchmark_duration,
|
||||||
tokenizer=tokenizer,
|
tokenizer=tokenizer,
|
||||||
|
|||||||
Reference in New Issue
Block a user