Track simulator fidelity experiment artifacts

This commit is contained in:
2026-07-19 15:31:09 +08:00
parent e0ea7e9961
commit 4c8d581a5b
115 changed files with 42355 additions and 0 deletions

View File

@@ -0,0 +1,200 @@
#!/usr/bin/env python3
"""Issue exact fixed-shape completion requests and record streaming latency."""
from __future__ import annotations
import argparse
import concurrent.futures
import http.client
import json
import statistics
import threading
import time
from pathlib import Path
from typing import Any
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, required=True)
parser.add_argument("--served-model", required=True)
parser.add_argument("--model-path", type=Path, required=True)
parser.add_argument("--input-tokens", type=int, default=2048)
parser.add_argument("--output-tokens", type=int, default=128)
parser.add_argument("--concurrency", type=int, required=True)
parser.add_argument("--requests", type=int, required=True)
parser.add_argument("--timeout-seconds", type=float, default=600.0)
parser.add_argument("--output", type=Path, required=True)
return parser.parse_args()
def run_request(
*,
request_index: int,
args: argparse.Namespace,
prompt_token_id: int,
start_barrier: threading.Barrier,
) -> dict[str, Any]:
body = {
"model": args.served_model,
"prompt": [prompt_token_id] * args.input_tokens,
"min_tokens": args.output_tokens,
"max_tokens": args.output_tokens,
"ignore_eos": True,
"temperature": 0,
"stream": True,
"stream_options": {"include_usage": True},
"return_token_ids": True,
}
encoded = json.dumps(body, separators=(",", ":")).encode()
connection = http.client.HTTPConnection(
args.host, args.port, timeout=args.timeout_seconds
)
start_barrier.wait()
started = time.perf_counter()
connection.request(
"POST",
"/v1/completions",
body=encoded,
headers={"Content-Type": "application/json"},
)
response = connection.getresponse()
if response.status != 200:
detail = response.read().decode(errors="replace")
raise RuntimeError(f"request {request_index} failed: HTTP {response.status}: {detail}")
first_token_at: float | None = None
last_token_at: float | None = None
streamed_token_count = 0
usage: dict[str, Any] | None = None
while True:
raw = response.readline()
if not raw:
break
line = raw.decode(errors="replace").strip()
if not line.startswith("data:"):
continue
data = line[5:].strip()
if data == "[DONE]":
break
payload = json.loads(data)
if payload.get("usage"):
usage = payload["usage"]
emitted = 0
for choice in payload.get("choices") or []:
token_ids = choice.get("token_ids") or []
if token_ids:
emitted += len(token_ids)
elif choice.get("text"):
emitted += 1
if emitted:
now = time.perf_counter()
if first_token_at is None:
first_token_at = now
last_token_at = now
streamed_token_count += emitted
finished = time.perf_counter()
connection.close()
if first_token_at is None or last_token_at is None or usage is None:
raise RuntimeError(
f"request {request_index} missing streaming token or usage metadata"
)
prompt_tokens = int(usage["prompt_tokens"])
completion_tokens = int(usage["completion_tokens"])
if prompt_tokens != args.input_tokens or completion_tokens != args.output_tokens:
raise RuntimeError(
f"request {request_index} usage mismatch: prompt={prompt_tokens}, "
f"completion={completion_tokens}"
)
ttft_ms = (first_token_at - started) * 1000.0
tpot_ms = (
(last_token_at - first_token_at) * 1000.0 / (completion_tokens - 1)
if completion_tokens > 1
else 0.0
)
return {
"request_index": request_index,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"streamed_token_count": streamed_token_count,
"ttft_ms": ttft_ms,
"tpot_ms": tpot_ms,
"e2e_ms": (finished - started) * 1000.0,
"ttft_slo_ms": 1000.0 + args.input_tokens / 8.0,
"tpot_slo_ms": 40.0,
"joint_slo_pass": ttft_ms <= 1000.0 + args.input_tokens / 8.0
and tpot_ms <= 40.0,
}
def percentile(values: list[float], fraction: float) -> float:
ordered = sorted(values)
index = min(len(ordered) - 1, max(0, round((len(ordered) - 1) * fraction)))
return ordered[index]
def main() -> None:
args = parse_args()
if args.concurrency <= 0 or args.requests < args.concurrency:
raise ValueError("requests must be at least concurrency, and both must be positive")
if args.input_tokens <= 0 or args.output_tokens <= 0:
raise ValueError("token lengths must be positive")
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)
candidate_ids = tokenizer.encode(" hello", add_special_tokens=False)
if not candidate_ids:
raise RuntimeError("tokenizer returned no prompt token id")
prompt_token_id = int(candidate_ids[0])
results: list[dict[str, Any]] = []
for batch_start in range(0, args.requests, args.concurrency):
batch_count = min(args.concurrency, args.requests - batch_start)
barrier = threading.Barrier(batch_count)
with concurrent.futures.ThreadPoolExecutor(max_workers=batch_count) as pool:
futures = [
pool.submit(
run_request,
request_index=batch_start + offset,
args=args,
prompt_token_id=prompt_token_id,
start_barrier=barrier,
)
for offset in range(batch_count)
]
results.extend(future.result() for future in futures)
ttfts = [float(row["ttft_ms"]) for row in results]
tpots = [float(row["tpot_ms"]) for row in results]
payload = {
"schema": "qwen235b-t0-smoke-v1",
"workload": {
"input_tokens": args.input_tokens,
"output_tokens": args.output_tokens,
"uniform_qps": None,
"prefix_caching": False,
"concurrency": args.concurrency,
"request_count": args.requests,
"prompt_token_id": prompt_token_id,
},
"summary": {
"completed_requests": len(results),
"joint_slo_pass_count": sum(bool(row["joint_slo_pass"]) for row in results),
"ttft_mean_ms": statistics.fmean(ttfts),
"ttft_p95_ms": percentile(ttfts, 0.95),
"tpot_mean_ms": statistics.fmean(tpots),
"tpot_p95_ms": percentile(tpots, 0.95),
},
"requests": sorted(results, key=lambda row: int(row["request_index"])),
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
print(json.dumps(payload["summary"], sort_keys=True), flush=True)
if __name__ == "__main__":
main()