File size: 8,897 Bytes
4d9b003 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | #!/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())
|