267 lines
9.2 KiB
Python
267 lines
9.2 KiB
Python
#!/usr/bin/env python3
|
|
"""Throughput benchmark for vLLM DSpark service (Qwen3-4B + dspark_qwen3_4b_block7)."""
|
|
|
|
import argparse
|
|
import asyncio
|
|
import json
|
|
import random
|
|
import time
|
|
from dataclasses import dataclass, field
|
|
from datetime import datetime
|
|
from typing import List
|
|
|
|
import aiohttp
|
|
import numpy as np
|
|
from transformers import AutoTokenizer
|
|
|
|
|
|
@dataclass
|
|
class RequestResult:
|
|
prompt_len: int = 0
|
|
output_len: int = 0
|
|
ttft_ms: float = 0.0
|
|
tpot_ms: float = 0.0
|
|
e2e_ms: float = 0.0
|
|
success: bool = False
|
|
error: str = ""
|
|
|
|
|
|
async def async_request(
|
|
session: aiohttp.ClientSession,
|
|
url: str,
|
|
model_name: str,
|
|
prompt: str,
|
|
max_tokens: int,
|
|
prompt_len: int,
|
|
result: RequestResult,
|
|
) -> None:
|
|
payload = {
|
|
"model": model_name,
|
|
"prompt": prompt,
|
|
"max_tokens": max_tokens,
|
|
"temperature": 0.0,
|
|
"stream": True,
|
|
"stream_options": {"include_usage": True},
|
|
}
|
|
result.prompt_len = prompt_len
|
|
start = time.perf_counter()
|
|
first_token_time = None
|
|
token_times: List[float] = []
|
|
output_text = ""
|
|
try:
|
|
async with session.post(url, json=payload) as resp:
|
|
resp.raise_for_status()
|
|
async for line in resp.content:
|
|
line = line.decode("utf-8").strip()
|
|
if not line or line.startswith(":"):
|
|
continue
|
|
if line.startswith("data: "):
|
|
data = line[len("data: "):]
|
|
if data == "[DONE]":
|
|
break
|
|
chunk = json.loads(data)
|
|
choices = chunk.get("choices", [])
|
|
if choices:
|
|
delta = choices[0].get("text", "")
|
|
if delta:
|
|
now = time.perf_counter()
|
|
token_times.append(now)
|
|
output_text += delta
|
|
if first_token_time is None:
|
|
first_token_time = now
|
|
end = time.perf_counter()
|
|
result.e2e_ms = (end - start) * 1000.0
|
|
if first_token_time is not None:
|
|
result.ttft_ms = (first_token_time - start) * 1000.0
|
|
if len(token_times) >= 2:
|
|
# TPOT: average time between consecutive tokens
|
|
intervals = [
|
|
token_times[i] - token_times[i - 1]
|
|
for i in range(1, len(token_times))
|
|
]
|
|
result.tpot_ms = sum(intervals) / len(intervals) * 1000.0
|
|
result.output_len = len(token_times)
|
|
result.success = True
|
|
except Exception as e:
|
|
result.e2e_ms = (time.perf_counter() - start) * 1000.0
|
|
result.error = str(e)
|
|
result.success = False
|
|
|
|
|
|
async def run_concurrency_benchmark(
|
|
url: str,
|
|
model_name: str,
|
|
prompts: List[str],
|
|
prompt_lens: List[int],
|
|
max_tokens: int,
|
|
concurrency: int,
|
|
num_prompts: int,
|
|
seed: int,
|
|
) -> dict:
|
|
rng = random.Random(seed)
|
|
indices = [rng.randrange(len(prompts)) for _ in range(num_prompts)]
|
|
|
|
semaphore = asyncio.Semaphore(concurrency)
|
|
results: List[RequestResult] = [RequestResult() for _ in range(num_prompts)]
|
|
|
|
async def bounded_request(idx: int, i: int):
|
|
async with semaphore:
|
|
async with aiohttp.ClientSession() as session:
|
|
await async_request(
|
|
session,
|
|
url,
|
|
model_name,
|
|
prompts[idx],
|
|
max_tokens,
|
|
prompt_lens[idx],
|
|
results[i],
|
|
)
|
|
|
|
start = time.perf_counter()
|
|
await asyncio.gather(*[bounded_request(idx, i) for i, idx in enumerate(indices)])
|
|
duration = time.perf_counter() - start
|
|
|
|
successes = [r for r in results if r.success]
|
|
failures = [r for r in results if not r.success]
|
|
if not successes:
|
|
return {
|
|
"concurrency": concurrency,
|
|
"num_prompts": num_prompts,
|
|
"duration_s": duration,
|
|
"success_count": 0,
|
|
"error": "all requests failed",
|
|
"errors": [r.error for r in failures[:5]],
|
|
}
|
|
|
|
total_in_tokens = sum(r.prompt_len for r in successes)
|
|
total_out_tokens = sum(r.output_len for r in successes)
|
|
total_tokens = total_in_tokens + total_out_tokens
|
|
|
|
ttfts = [r.ttft_ms for r in successes]
|
|
tpots = [r.tpot_ms for r in successes]
|
|
e2es = [r.e2e_ms for r in successes]
|
|
|
|
return {
|
|
"concurrency": concurrency,
|
|
"num_prompts": num_prompts,
|
|
"duration_s": duration,
|
|
"success_count": len(successes),
|
|
"fail_count": len(failures),
|
|
"input_tokens": total_in_tokens,
|
|
"output_tokens": total_out_tokens,
|
|
"total_tokens": total_tokens,
|
|
"request_throughput": len(successes) / duration,
|
|
"input_throughput": total_in_tokens / duration,
|
|
"output_throughput": total_out_tokens / duration,
|
|
"total_throughput": total_tokens / duration,
|
|
"mean_ttft_ms": float(np.mean(ttfts)),
|
|
"p50_ttft_ms": float(np.percentile(ttfts, 50)),
|
|
"p99_ttft_ms": float(np.percentile(ttfts, 99)),
|
|
"mean_tpot_ms": float(np.mean(tpots)),
|
|
"p50_tpot_ms": float(np.percentile(tpots, 50)),
|
|
"p99_tpot_ms": float(np.percentile(tpots, 99)),
|
|
"mean_e2e_ms": float(np.mean(e2es)),
|
|
"p50_e2e_ms": float(np.percentile(e2es, 50)),
|
|
"p99_e2e_ms": float(np.percentile(e2es, 99)),
|
|
"errors": [r.error for r in failures[:5]],
|
|
}
|
|
|
|
|
|
def load_sharegpt_prompts(path: str, tokenizer, max_samples: int = 2000):
|
|
with open(path, "r") as f:
|
|
data = json.load(f)
|
|
prompts = []
|
|
lens = []
|
|
for item in data:
|
|
conv = item.get("conversations", [])
|
|
for turn in conv:
|
|
if turn.get("from") == "human":
|
|
prompt = turn.get("value", "")
|
|
if prompt:
|
|
prompts.append(prompt)
|
|
lens.append(len(tokenizer.encode(prompt)))
|
|
break
|
|
if len(prompts) >= max_samples:
|
|
break
|
|
return prompts, lens
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--host", default="127.0.0.1")
|
|
parser.add_argument("--port", type=int, default=30003)
|
|
parser.add_argument("--model", default="/data/models/Qwen3-4B")
|
|
parser.add_argument("--dataset", default="/data/user1/yy/datasets/ShareGPT_filtered_chat.json")
|
|
parser.add_argument("--tokenizer", default="/data/models/Qwen3-4B")
|
|
parser.add_argument("--max-tokens", type=int, default=256)
|
|
parser.add_argument("--num-prompts", type=int, default=500)
|
|
parser.add_argument("--concurrency", type=int, nargs="+", default=[1, 4, 16, 64, 128])
|
|
parser.add_argument("--seed", type=int, default=42)
|
|
parser.add_argument("--output-dir", default="/data/user1/yy/bench_results")
|
|
args = parser.parse_args()
|
|
|
|
url = f"http://{args.host}:{args.port}/v1/completions"
|
|
print(f"Loading tokenizer from {args.tokenizer} ...")
|
|
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer)
|
|
print(f"Loading dataset from {args.dataset} ...")
|
|
prompts, prompt_lens = load_sharegpt_prompts(args.dataset, tokenizer)
|
|
print(f"Loaded {len(prompts)} prompts, prompt lens: min={min(prompt_lens)}, max={max(prompt_lens)}, mean={sum(prompt_lens)/len(prompt_lens):.1f}")
|
|
|
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
|
result_dir = f"{args.output_dir}/vllm_dspark_qwen3_{timestamp}"
|
|
import os
|
|
os.makedirs(result_dir, exist_ok=True)
|
|
summary_file = f"{result_dir}/summary.json"
|
|
|
|
summary = {
|
|
"model": args.model,
|
|
"draft_model": "deepseek-ai/dspark_qwen3_4b_block7",
|
|
"url": url,
|
|
"max_tokens": args.max_tokens,
|
|
"num_prompts": args.num_prompts,
|
|
"dataset": args.dataset,
|
|
"timestamp": timestamp,
|
|
"results": [],
|
|
}
|
|
|
|
print("\nStarting benchmark sweeps...")
|
|
print(f"{'Concurrency':>12} {'Req/s':>10} {'In tok/s':>12} {'Out tok/s':>12} {'Total tok/s':>13} {'TTFT(ms)':>10} {'TPOT(ms)':>10} {'E2E(ms)':>10}")
|
|
print("-" * 100)
|
|
|
|
for c in args.concurrency:
|
|
print(f"Running concurrency={c} ...", flush=True)
|
|
result = asyncio.run(
|
|
run_concurrency_benchmark(
|
|
url,
|
|
args.model,
|
|
prompts,
|
|
prompt_lens,
|
|
args.max_tokens,
|
|
c,
|
|
args.num_prompts,
|
|
args.seed,
|
|
)
|
|
)
|
|
summary["results"].append(result)
|
|
with open(summary_file, "w") as f:
|
|
json.dump(summary, f, indent=2)
|
|
if result.get("success_count", 0) == 0:
|
|
print(f"{c:>12} FAILED: {result.get('error', 'unknown')}")
|
|
continue
|
|
print(
|
|
f"{c:>12} "
|
|
f"{result['request_throughput']:>10.2f} "
|
|
f"{result['input_throughput']:>12.2f} "
|
|
f"{result['output_throughput']:>12.2f} "
|
|
f"{result['total_throughput']:>13.2f} "
|
|
f"{result['mean_ttft_ms']:>10.1f} "
|
|
f"{result['mean_tpot_ms']:>10.1f} "
|
|
f"{result['mean_e2e_ms']:>10.1f}"
|
|
)
|
|
|
|
print(f"\nSummary saved to {summary_file}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|