#!/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 /diffusion-planner-p150/tt_kernel_manifest.json [--profile ] \ --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 ``..reference.json``, ``..reference.json``, ``.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())