changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
9.01 kB
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Goldens of the fp32 CPU reference (no device): per-module taps and final outputs per scene.
# research venv (torch + onnx; onnxruntime only for --ort), from the bundle root:
PYTHONPATH=code python code/scripts/ref_golden.py --threads 4 [--ort] \
[--full-dir ../../research/diffusion-planner/goldens] [--scenes kashiwanoha_dense straight_road ...]
Scenes: the shipped samples (``code/tt_diffusion_planner/samples/*.npz``: small goldens -> ``tests/goldens``, and the
stored ``/predict`` body of the reference -> ``samples/<stem>.reference.json``) and, when the workspace research
directory is present, the seven ORT golden scenes of ``research/diffusion-planner/ort/golden_<scene>.npz`` and any
public-dataset scene in ``research/diffusion-planner/public_data/*.npz`` (full goldens only). ``--ort`` also runs ONNX
Runtime on every scene and writes ``<full-dir>/<scene>.ort_agreement.json`` (reference vs ORT on the deployed ONNX).
"""
from __future__ import annotations
import argparse
import hashlib
import json
import sys
import time
from pathlib import Path
import numpy as np
CODE = Path(__file__).resolve().parents[1]
if str(CODE) not in sys.path:
sys.path.insert(0, str(CODE))
from tt_diffusion_planner.reference import config as C # noqa: E402
from tt_diffusion_planner.reference.goldens import write_scene # noqa: E402
from tt_diffusion_planner.reference.pipeline import ReferencePlanner # noqa: E402
from tt_diffusion_planner.ttaw.metrics import pcc # noqa: E402
PKG = CODE / "tt_diffusion_planner"
SAMPLES = PKG / "samples"
SMALL_DIR = PKG / "tests" / "goldens"
RESEARCH = CODE.parents[2] / "research" / "diffusion-planner"
def public_ids(dataset: str, which: str):
"""Instant ids of ``public_data/inputs/<dataset>``: ``all``, or ``core`` (the ids the public-data goldens keep
solver iterates for: ``public_data/goldens/<dataset>/core_ids.json``)."""
root = RESEARCH / "public_data"
ids = sorted(p.stem for p in (root / "inputs" / dataset).glob("*.npz"))
if which == "core":
core = root / "goldens" / dataset / "core_ids.json"
keep = set(json.loads(core.read_text()).get("core_ids_with_denoising_steps", [])) if core.is_file() else set()
ids = [i for i in ids if i in keep]
return ids
def scene_sources(selected, public=None, public_which="core"):
"""``{scene: (raw source, kind)}``: shipped samples first, then research ORT scenes, then (``--public``) the
public-data instants of one dataset (nuScenes-derived: CC BY-NC-SA, local goldens only)."""
out = {}
if public is None:
for p in sorted(SAMPLES.glob("*.npz")):
out[p.stem] = (p, "sample")
for p in sorted((RESEARCH / "ort").glob("golden_*.npz")):
out.setdefault(p.stem[len("golden_"):], (p, "research"))
else:
for i in public_ids(public, public_which):
out[f"public/{public}/{i}"] = (RESEARCH / "public_data" / "inputs" / public / f"{i}.npz", "public")
if selected:
missing = sorted(set(selected) - set(out))
if missing:
raise SystemExit(f"unknown scenes {missing}; have {sorted(out)}")
out = {k: v for k, v in out.items() if k in selected}
return out
def load_raw(path: Path, kind: str):
with np.load(path, allow_pickle=False) as z:
if kind in ("research", "public"): # dp_reference.py / nuscenes_dp.py layout: raw/<name>
return {k: np.asarray(z["raw/" + k], np.float32) for k in C.INPUT_NAMES}
return {k: np.asarray(z[k], np.float32) for k in C.INPUT_NAMES}
def ort_agreement(weights_dir: Path, raw, g, threads: int):
from tt_diffusion_planner.reference.ort import OrtPlanner
o = OrtPlanner(weights_dir, threads=threads).run(raw)
rows = g["dec.rows"]
enc_rows = np.flatnonzero(g["host.token_valid"])
return {"encoding_pcc_valid": pcc(g["enc.encoding"][enc_rows], o.encoding[0][enc_rows]),
"encoding_max_abs": float(np.abs(g["enc.encoding"] - o.encoding[0]).max()),
"final_x0_pcc_valid": pcc(g["final_x0"][rows], o.final_x0[0][rows]),
"final_x0_max_abs_valid": float(np.abs(g["final_x0"][rows] - o.final_x0[0][rows]).max()),
"logit_max_abs": float(np.abs(g["turn.logit"] - o.logit[0]).max()),
"turn_logit_ref": g["turn.logit"].tolist(), "turn_logit_ort": o.logit[0].tolist()}
def stored_ort_agreement(g, golden: Path):
"""Reference vs a stored ORT golden in the ``dp_reference.py`` layout (``encoding``, ``final_x_normalized``,
``logit_multi``): the public-data goldens of ``research/diffusion-planner/public_data/goldens/<dataset>``."""
with np.load(golden, allow_pickle=False) as z:
enc, fx, lg = z["encoding"][0], z["final_x_normalized"][0], z["logit_multi"][0]
rows = g["dec.rows"]
tok = np.flatnonzero(g["host.token_valid"])
return {"ort_golden": str(golden),
"encoding_pcc_valid": pcc(g["enc.encoding"][tok], enc[tok]),
"encoding_max_abs": float(np.abs(g["enc.encoding"][tok] - enc[tok]).max()),
"final_x0_pcc_valid": pcc(g["final_x0"][rows], fx[rows]),
"final_x0_max_abs_valid": float(np.abs(g["final_x0"][rows] - fx[rows]).max()),
"logit_max_abs": float(np.abs(g["turn.logit"] - lg).max())}
def main(argv=None) -> int:
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--weights-dir", default=None)
ap.add_argument("--full-dir", type=Path, default=RESEARCH / "goldens",
help="where the large per-scene goldens go (never inside the bundle)")
ap.add_argument("--scenes", nargs="*", default=None)
ap.add_argument("--threads", type=int, default=4)
ap.add_argument("--ort", action="store_true", help="also compare with ONNX Runtime (research venv)")
ap.add_argument("--no-reference-json", action="store_true")
ap.add_argument("--public", default=None, help="a public_data/inputs/<dataset> (e.g. nuscenes): its instants "
"instead of the samples / research scenes, as lite goldens")
ap.add_argument("--public-ids", default="core", choices=["core", "all"])
a = ap.parse_args(argv)
if a.full_dir.resolve().is_relative_to(CODE.parent.resolve()):
raise SystemExit("--full-dir must be outside the bundle: the full goldens are 10-25 MB per scene")
ref = ReferencePlanner(a.weights_dir, threads=a.threads)
report, seen = {}, {}
for scene, (path, kind) in scene_sources(a.scenes, a.public, a.public_ids).items():
t0 = time.perf_counter()
raw = load_raw(path, kind)
digest = hashlib.sha256(b"".join(raw[k].tobytes() for k in C.INPUT_NAMES)).hexdigest()
if digest in seen: # e.g. research golden_straight == the shipped straight_road sample
report[scene] = {"same_inputs_as": seen[digest]}
print(scene, "skipped: same inputs as", seen[digest], flush=True)
continue
seen[digest] = scene
small = SMALL_DIR if kind == "sample" else None
r = write_scene(ref, raw, scene, a.full_dir, small, meta={"source": str(path), "kind": kind},
lite=(kind == "public"))
entry = {"paths": r["paths"], "valid_counts": r["info"]["valid_counts"]}
if kind == "sample" and not a.no_reference_json:
body = ref(inputs=raw).to_dict()
body["timing_ms"] = {}
ref_path = SAMPLES / f"{scene}.reference.json"
ref_path.write_text(json.dumps(body, indent=1) + "\n")
entry["reference_json"] = str(ref_path)
stored = None
if kind == "public":
stored = RESEARCH / "public_data" / "goldens" / a.public / f"golden_{path.stem}.npz"
if stored is not None and stored.is_file(): # the dataset work already ran ORT on this instant
agree = stored_ort_agreement(r["goldens"], stored)
(a.full_dir / f"{scene}.ort_agreement.json").write_text(json.dumps(agree, indent=1) + "\n")
entry["ort"] = {k: v for k, v in agree.items() if k != "ort_golden"}
elif a.ort:
agree = ort_agreement(ref.weights.path, raw, r["goldens"], a.threads)
a.full_dir.mkdir(parents=True, exist_ok=True)
(a.full_dir / f"{scene}.ort_agreement.json").write_text(json.dumps(agree, indent=1) + "\n")
entry["ort"] = {k: v for k, v in agree.items() if not k.startswith("turn_logit")}
entry["seconds"] = round(time.perf_counter() - t0, 1)
report[scene] = entry
print(scene, json.dumps(entry), flush=True)
a.full_dir.mkdir(parents=True, exist_ok=True)
index = a.full_dir / "index.json"
merged = json.loads(index.read_text()) if index.is_file() else {}
merged.update(report)
index.write_text(json.dumps(merged, indent=1, sort_keys=True) + "\n")
return 0
if __name__ == "__main__":
sys.exit(main())