# SPDX-License-Identifier: Apache-2.0 """Where does the end-to-end error of a scene come from? (development tool; device, run under ``bin/devrun``) bin/devrun -t 1200 -- python code/scripts/split_error.py --public scene-0103_kf14 --scenes straight_road For each scene the ego / neighbour displacement vs the fp32 CPU reference of: - ``device``: the full device plan; - ``enc_only``: the device encoding (``encoder_taps``) + the CPU fp32 decoder and solver (encoder error alone); - ``dec_only``: the CPU encoding + the device decoder (``decode_once`` replays) driven by the host solver (decoder error alone, including its compounding over the 11 evaluations). """ from __future__ import annotations import argparse import json import sys from pathlib import Path import numpy as np HERE = Path(__file__).resolve() sys.path.insert(0, str(HERE.parents[1])) from tt_diffusion_planner.host import pipeline as hp # noqa: E402 from tt_diffusion_planner.host.solver import apply_prefix_constraint, dpm_solver_sample # noqa: E402 from tt_diffusion_planner.reference import config as C # noqa: E402 RES = HERE.parents[4] / "research" / "diffusion-planner" def load_raw(scene: str, public: bool): if public: with np.load(RES / "public_data" / "inputs" / "nuscenes" / f"{scene}.npz", allow_pickle=False) as z: return {k: np.array(z[f"raw/{k}"]) for k in C.INPUT_NAMES} with np.load(RES / "goldens" / f"{scene}.npz", allow_pickle=False) as z: return {k: np.array(z[f"in.{k}"]) for k in C.INPUT_NAMES} def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--public", nargs="*", default=[]) ap.add_argument("--scenes", nargs="*", default=[]) ap.add_argument("--json", default=None) for opt in ("ln-fp32", "hidden-fp32", "split", "attn-fp32-acc", "attn-matmul"): ap.add_argument(f"--{opt}", default=None, help="module globs (comma-separated); default: the knob") args = ap.parse_args() from tt_diffusion_planner.tt.config import globs opts = {k: globs(getattr(args, k)) for k in ("ln_fp32", "hidden_fp32", "split", "attn_fp32_acc", "attn_matmul") if getattr(args, k) is not None} import torch torch.set_num_threads(4) from tt_diffusion_planner.device import close_device, open_device from tt_diffusion_planner.reference.model import Decoder, Encoder, torch_params from tt_diffusion_planner.reference.weights import find_weights_dir, load_weights from tt_diffusion_planner.tt.model import TtDiffusionPlanner w = load_weights(find_weights_dir()) params = {k: v[3] for k, v in hp.RUNTIME_PARAMS.items()} P = torch_params(w.params) enc_ref, dec_ref = Encoder(P), Decoder(P) dev = open_device(allow_fallback=False) report = {} try: tt = TtDiffusionPlanner(dev, w, debug=True, **opts) report["options"] = tt.build.options() tt.capture() for scene, public in [(s, True) for s in args.public] + [(s, False) for s in args.scenes]: prep = hp.prepare(load_raw(scene, public), w.normalization.observation) cs = prep.decoder.current_states def solve(model_fn): res = dpm_solver_sample(prep.x_T, model_fn, lambda x: apply_prefix_constraint(x, cs)) return res.final_x def cpu_plan(encoding): kv = dec_ref.cross_kv(torch.from_numpy(np.asarray(encoding, np.float32))) with torch.no_grad(): return solve(lambda x, t: dec_ref.forward(x, t, kv, prep.decoder.agent_valid).numpy()) def out(final_x0): o = hp.make_output(final_x0, np.zeros(5, np.float32), prep, w.normalization, params) return o.poses[:, :2], o.predicted_agents[..., :2] with torch.no_grad(): enc = enc_ref.forward(prep.features).numpy() ref_ego, ref_nb = out(cpu_plan(enc)) dev_enc = tt.encoder_taps(prep)["enc.encoding"] cases = { "device": tt.forward(prep)["final_x0"], "enc_only": cpu_plan(dev_enc), "dec_only": solve(lambda x, t: _decode(tt, prep, x, t, enc)), } res = {} for name, fx in cases.items(): ego, nb = out(fx) d = np.hypot(*(ego - ref_ego).T) r = {"ego_max_m": round(float(d.max()), 4), "ego_mean_m": round(float(d.mean()), 4)} if nb.shape[0]: per_agent = np.hypot(*(nb - ref_nb).transpose(2, 0, 1)).max(1) r["nb_median_max_m"] = round(float(np.median(per_agent)), 4) res[name] = r res["encoding_pcc"] = float(np.corrcoef(dev_enc[prep.features.token_valid].ravel(), enc[prep.features.token_valid].ravel())[0, 1]) report[scene] = res print(scene, json.dumps(res), flush=True) tt.release() finally: close_device(dev) if args.json: Path(args.json).write_text(json.dumps(report, indent=1) + "\n") return 0 def _decode(tt, prep, x, t, enc): out = tt.decode_once(prep, x, float(t), encoding=enc) out[:, 0] = 0.0 # the t = 0 slot (masked on the device) is overwritten by the prefix constraint anyway return out if __name__ == "__main__": sys.exit(main())