GR00T-N1.7-3B-p150 / code /scripts /bench_http.py
changh95's picture
Add files using upload-large-folder tool
63804fd verified
Raw History Blame Contribute Delete
6.39 kB
#!/usr/bin/env python3
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""Latency of a running GR00T tt-dit-server over HTTP: posts the shipped demo observation N times and reports the
server-side ``timing_ms`` (``device`` = the four Metal trace replays incl. input upload and readback, ``total`` = the
whole handler) and the client wall time. Standard library only (any Python >= 3.8).
::
python scripts/bench_http.py --url http://127.0.0.1:20000 [--n 50] [--warmup 5] [--out bench.json]
The request uses the server's default seed (the deployed policy's fixed noise), so the payload is the demo frames (one PNG per camera and history slot) plus
the raw state; every request has the same shape, which is what a real control loop sends.
"""
from __future__ import annotations
import argparse
import base64
import json
import statistics
import sys
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any, Dict, List, Optional
DEFAULT_DEMO_ROOT = Path(__file__).resolve().parents[1] / "gr00t_p150" / "demo"
def _get(url: str, timeout: float = 30.0) -> Dict[str, Any]:
with urllib.request.urlopen(url, timeout=timeout) as r:
return json.loads(r.read().decode())
def _post(url: str, payload: Dict[str, Any], timeout: float = 600.0) -> Dict[str, Any]:
data = json.dumps(payload).encode()
req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"}, method="POST")
try:
with urllib.request.urlopen(req, timeout=timeout) as r:
return json.loads(r.read().decode())
except urllib.error.HTTPError as e:
raise SystemExit(f"FAIL HTTP {e.code}: {e.read().decode(errors='replace')[:400]}") from None
def wait_ready(url: str, timeout: float) -> None:
deadline = time.time() + timeout
last: Any = None
while time.time() < deadline:
try:
h = _get(f"{url}/health", timeout=10)
last = h
if h.get("status") == "ok":
return
except Exception as e: # noqa: BLE001
last = f"{type(e).__name__}: {e}"
time.sleep(5)
raise SystemExit(f"FAIL server at {url} not ready after {timeout:.0f}s (last: {last})")
def demo_request(demo_dir: Path) -> Dict[str, Any]:
"""The shipped demo observation as a ``/predict`` body (seed path: no explicit noise)."""
with open(demo_dir / "observation.json") as fh:
obs = json.load(fh)
images: Dict[str, List[str]] = {}
for cam, files in obs["cameras"].items():
images[cam] = [base64.b64encode((demo_dir / f).read_bytes()).decode("ascii") for f in files]
return {
"images": images,
"state": obs["state"],
"instruction": obs["instruction"],
"embodiment": obs["embodiment"],
"state_dtype": obs["state_dtype"],
}
def percentile(xs: List[float], q: float) -> float:
if not xs:
raise ValueError("empty sample")
s = sorted(xs)
k = (len(s) - 1) * q
lo, hi = int(k), min(int(k) + 1, len(s) - 1)
return s[lo] + (s[hi] - s[lo]) * (k - lo)
def summarize(xs: List[float]) -> Dict[str, float]:
return {
"median": statistics.median(xs),
"p90": percentile(xs, 0.90),
"min": min(xs),
"max": max(xs),
"mean": statistics.fmean(xs),
"n": float(len(xs)),
}
def main(argv: Optional[List[str]] = None) -> int:
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--url", default="http://127.0.0.1:20000")
ap.add_argument("--n", type=int, default=50, help="timed requests (default 50)")
ap.add_argument("--warmup", type=int, default=5, help="untimed requests first (default 5)")
ap.add_argument("--wait", type=float, default=1800.0, help="seconds to wait for /health == ok")
ap.add_argument("--version", default=None, choices=["n15", "n16", "n17"], help="default: what /info reports")
ap.add_argument("--demo-dir", type=Path, default=None, help="default: gr00t_p150/demo/<version>")
ap.add_argument("--out", type=Path, default=None, help="write every sample + the summary as JSON")
args = ap.parse_args(argv)
if args.n < 1 or args.warmup < 0:
raise SystemExit("--n must be >= 1 and --warmup >= 0")
wait_ready(args.url, args.wait)
info = _get(f"{args.url}/info")
version = args.version or str(info.get("version"))
demo_dir = args.demo_dir or (DEFAULT_DEMO_ROOT / version)
body = demo_request(demo_dir)
predict = f"{args.url}/predict"
for _ in range(args.warmup):
_post(predict, body)
samples: List[Dict[str, float]] = []
first: Optional[Dict[str, Any]] = None
for _ in range(args.n):
t0 = time.perf_counter()
resp = _post(predict, body)
wall = (time.perf_counter() - t0) * 1000.0
if first is None:
first = resp
t = resp["timing_ms"]
samples.append({"wall": wall, **{k: float(v) for k, v in t.items()}})
keys = sorted(samples[0])
summary = {k: summarize([s[k] for s in samples]) for k in keys}
model = info.get("model", "?")
dev, tot, wall = summary["device"], summary["total"], summary["wall"]
line = (
f"BENCH {model} {version}: n={args.n} warmup={args.warmup} "
f"device_ms median={dev['median']:.2f} p90={dev['p90']:.2f} min={dev['min']:.2f} max={dev['max']:.2f} | "
f"total_ms median={tot['median']:.2f} p90={tot['p90']:.2f} | "
f"client_wall_ms median={wall['median']:.2f} p90={wall['p90']:.2f}"
)
result = {
"url": args.url,
"model": model,
"version": version,
"n": args.n,
"warmup": args.warmup,
"request": {"seed": "server default", "cameras": sorted(body["images"])},
"info_warmup_latency_ms": info.get("warmup_latency_ms"),
"summary": summary,
"samples": samples,
"first_response_meta": {
k: v for k, v in (first or {}).items() if k not in ("actions", "action_pred_normalized")
},
"line": line,
}
if args.out:
args.out.parent.mkdir(parents=True, exist_ok=True)
with open(args.out, "w") as fh:
json.dump(result, fh, indent=1)
print(line)
return 0
if __name__ == "__main__":
sys.exit(main())