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())