changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
8.9 kB
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Smoke test of a running diffusion-planner-p150 server (standard library only; any Python 3.9+, no numpy).
python3 code/tt_diffusion_planner/server/smoke_test.py --url http://127.0.0.1:20000 --wait 1800
# a served container package, one serve profile (what code/scripts/container_smoke.sh runs):
python3 code/tt_diffusion_planner/server/smoke_test.py --url http://127.0.0.1:20000 --wait 600 \
--manifest <out>/diffusion-planner-p150/tt_kernel_manifest.json [--profile <name>] \
--out /tmp/diffusion-planner.json
Checks (all failures are collected; prints ONE line ``PASS ...`` / ``FAIL ...``, exit code 0 / 1):
* ``/health`` reports ``ok`` (waiting up to ``--wait`` seconds), ``/info`` and ``/v1/models`` answer;
* ``/info`` reports ETH dispatch and the 12x10 grid, the p150 target of every published number (PLAN.md 0.3
item 6, D14): WORKER dispatch (e.g. an ETH open that fell back) or any other grid FAILS;
* with ``--manifest``: ``/info`` runs what the staged package pins for the serve profile
(``DIFFUSION_PLANNER_DISPATCH``,
``DIFFUSION_PLANNER_NUM_CQS``, ``DIFFUSION_PLANNER_VARIANT``, the weights revision);
* ``POST /predict`` of the shipped sample (the planner tensors as ``inputs``) answers 200 with the documented fields
and passes this model's output gates (:func:`output_gates`: 80 finite trajectory rows of 7 columns, a valid
turn-indicator command, the predicted neighbour paths);
* the served output agrees with the stored CPU-reference output of the sample within :data:`REFERENCE_GATES`:
``--reference``, else the first of ``<sample stem>.<profile>.reference.json``,
``<stem>.<variant>.reference.json``, ``<stem>.reference.json`` next to the sample (skipped when none exists);
* a malformed request answers 400 (not 500).
The checks themselves are the vendored ``ttaw/server/client.py`` and ``ttaw/server/smoke.py`` (stdlib only, loaded
by path); this file holds the model-specific parts: the sample request, the output gates and their thresholds.
"""
from __future__ import annotations
import argparse
import json
import math
import sys
import time
from pathlib import Path
from typing import Any, Dict, List, Optional
HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE.parent / "ttaw" / "server"))
import client as ttaw_client # noqa: E402 (stdlib-only modules of the vendored ttaw)
import smoke as ttaw_smoke # noqa: E402
MODEL = "diffusion-planner-p150"
ENV_PREFIX = "DIFFUSION_PLANNER"
EXPECT_DISPATCH, EXPECT_GRID = "eth", "12x10" # the p150 target; deliberately not a command-line option
DEFAULT_INPUT = HERE.parent / "samples" / "kashiwanoha_dense.npz"
DEFAULT_EXPECT = "" # unused for a planner (no detections); kept for the shared command line
# Agreement with the stored CPU reference (names of ttaw_smoke.DEFAULT_GATES): the ego trajectory's average / final
# displacement (the predicted_agents array is reported, not gated). Keep them in line with tests/test_e2e_device.py
# (ego mean error <= 0.3 m; the max error <= 1.0 m gate needs the whole array and lives in the device test).
REFERENCE_GATES: Dict[str, Optional[float]] = {"max_ade": 0.3, "max_fde": 1.0}
REQUIRED_KEYS = ("model", "frame_id", "timing_ms", "num_poses", "columns", "trajectory", "turn_indicator",
"predicted_agents")
TRAJECTORY_COLUMNS = ["x", "y", "yaw", "cos", "sin", "velocity", "acceleration"]
def build_sample_request(path: Path) -> Dict[str, Any]:
"""The ``/predict`` body for the sample: the planner tensors ``.npz`` as ``inputs``."""
if not Path(path).is_file():
raise FileNotFoundError(str(path))
return ttaw_client.build_request(inputs=str(path))
def output_gates(body: Dict[str, Any], expect: str) -> List[str]:
"""This model's plausibility gates on a 200 response: 80 trajectory rows of the documented 7 finite columns,
a turn-indicator command in 0..3 with 5 logits, and an encoded predicted-agents array of 80 x 5 per agent."""
fails: List[str] = []
traj = body.get("trajectory") or []
if body.get("columns") != TRAJECTORY_COLUMNS:
fails.append(f"columns {body.get('columns')} != {TRAJECTORY_COLUMNS}")
if len(traj) != 80 or body.get("num_poses") != 80:
fails.append(f"{len(traj)} trajectory rows, expected 80")
if not all(len(row) == 7 and all(math.isfinite(v) for v in row) for row in traj):
fails.append("non-finite or malformed trajectory row")
turn = body.get("turn_indicator") or {}
if turn.get("command") not in (0, 1, 2, 3) or len(turn.get("logits") or []) != 5:
fails.append(f"turn_indicator {turn}")
agents = body.get("predicted_agents") or {}
shape = agents.get("shape") or []
if len(shape) != 3 or shape[1:] != [80, 5]:
fails.append(f"predicted_agents shape {shape}")
return fails
def main(argv: Optional[List[str]] = None) -> int:
ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
ap.add_argument("--url", default="http://127.0.0.1:20000")
ap.add_argument("--input", type=Path, default=DEFAULT_INPUT, help="sample to POST (default: the shipped one)")
ap.add_argument("--expect", default=DEFAULT_EXPECT, help="unused by this planner (kept for the shared CLI)")
ap.add_argument("--reference", type=Path, help="stored CPU-reference /predict body of --input "
"(default: looked up next to it)")
ap.add_argument("--manifest", type=Path, help="staged tt_kernel_manifest.json: /info must run what it pins")
ap.add_argument("--profile", help="serve profile being smoke-tested (default: the package's default)")
ap.add_argument("--wait", type=float, default=0.0, help="seconds to wait for /health == ok")
ap.add_argument("--out", type=Path, help="write the /predict response here")
a = ap.parse_args(argv)
base = a.url.rstrip("/")
health = ttaw_client.wait_ready(base, wait_s=a.wait)
if health.get("status") != "ok":
print(ttaw_client.smoke_line(MODEL, f"/health status={health.get('status')!r}",
[f"not ready after {a.wait:.0f} s (error: {health.get('error')})"]))
return 1
info, fails = ttaw_client.check_service(base, expect_dispatch=EXPECT_DISPATCH, expect_grid=EXPECT_GRID)
if not info:
print(ttaw_client.smoke_line(MODEL, "/info unreachable", fails))
return 1
profile = a.profile
if a.manifest:
try:
pinned = ttaw_smoke.pinned_config(a.manifest, a.profile)
except (OSError, ValueError) as e:
fails.append(f"manifest: {e}")
else:
profile = pinned["profile"]
fails += ttaw_smoke.check_pinned(info, pinned, ENV_PREFIX)
code, body = 0, {}
t0 = time.perf_counter()
try:
code, body = ttaw_client.post(base, build_sample_request(a.input))
except (OSError, ValueError) as e: # unreadable sample, unknown suffix, server gone
fails.append(f"/predict: {type(e).__name__}: {e}")
rtt_ms = (time.perf_counter() - t0) * 1e3
if a.out:
a.out.write_text(json.dumps(body, indent=1))
ref_summary = "none"
if code and code != 200:
fails.append(f"/predict HTTP {code}: {str(body)[:300]}")
elif code == 200:
fails += [f"missing key {k!r}" for k in REQUIRED_KEYS if k not in body]
fails += output_gates(body, a.expect)
ref_path = a.reference or ttaw_smoke.find_reference(a.input, profile, info.get("variant"))
if ref_path:
try:
metrics, more = ttaw_smoke.compare_with_reference(body, ttaw_smoke.load_json(ref_path),
gates=REFERENCE_GATES)
except (OSError, ValueError) as e:
metrics, more = {}, [f"reference {ref_path}: {e}"]
fails += more
ref_summary = f"{ref_path.name} ({ttaw_smoke.describe_metrics(metrics)})"
try:
bad = ttaw_client.check_bad_request(base)
except OSError as e:
bad = f"malformed-request check: {type(e).__name__}: {e}"
if bad:
fails.append(bad)
device, timing = info.get("device") or {}, body.get("timing_ms") or {}
turn = (body.get("turn_indicator") or {}).get("command_name")
summary = (f"profile={profile or '-'} variant={info.get('variant')} dispatch={device.get('dispatch')} "
f"grid={device.get('grid')} cqs={device.get('num_command_queues')} n={body.get('num_poses')} "
f"turn={turn} "
f"reference={ref_summary} device_ms={timing.get('device')} total_ms={timing.get('total')} "
f"rtt_ms={rtt_ms:.1f}")
print(ttaw_client.smoke_line(MODEL, summary, fails))
return 1 if fails else 0
if __name__ == "__main__":
sys.exit(main())