Download code/tt_diffusion_planner/server/smoke_test.py from changh95/diffusion-planner-p150: direct link, hf CLI and curl.
- Browser
- Download file 8.9 kB
-
https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/server/smoke_test.py
- Command line
-
hf download hf://changh95/diffusion-planner-p150/code/tt_diffusion_planner/server/smoke_test.py
-
curl -L -o smoke_test.py https://huggingface.co/changh95/diffusion-planner-p150/resolve/main/code/tt_diffusion_planner/server/smoke_test.py
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()) | |