diffusion-planner-p150 / code /scripts /bringup_device.py
changh95's picture
tt-model push diffusion-planner-p150 (container)
4d9b003 verified
Raw History Blame Contribute Delete
9.23 kB
# SPDX-License-Identifier: Apache-2.0
"""Device bring-up of the planner graph (development tool; needs the p150, run through ``bin/devrun``).
bin/devrun -t 1800 -- python code/scripts/bringup_device.py [--scenes kashiwanoha_dense straight_road]
[--no-capture] [--precision "dec.*=HiFi2+fp32"] [--json logs/diffusion-planner/bringup.json]
Protocol (PLAN.md 4.4): build ``TtDiffusionPlanner(debug=True)``; per scene run every variant EAGERLY twice
(``encoder_taps``, ``decode_once`` at evaluations 0 / 5 / 10, ``plan``), check the two eager runs are bit-identical and
compare them with the research goldens (``research/diffusion-planner/goldens/<scene>.npz``, the fp32 CPU reference:
PCC on valid rows, max abs; final x0 on valid agents; logits); then capture all variants (strict, no program
compiled after the capture) and check replay == eager bit for bit on the same inputs. Prints a table and writes a
JSON report. This is a diagnostic: the frozen gates live in ``tests/test_pcc_device.py`` / ``test_e2e_device.py``.
"""
from __future__ import annotations
import argparse
import json
import sys
import time
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.reference import config as C # noqa: E402
from tt_diffusion_planner.ttaw.metrics import error_stats # noqa: E402
GOLDENS = HERE.parents[4] / "research" / "diffusion-planner" / "goldens"
def stats(dev, ref):
s = error_stats(np.asarray(dev, np.float64), np.asarray(ref, np.float64))
return {"pcc": round(s["pcc"], 7), "max_abs": float(f"{s['max_abs']:.4g}"), "rel_l2": float(f"{s['rel_l2']:.4g}")}
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
ap.add_argument("--scenes", nargs="+", default=["kashiwanoha_dense", "straight_road"])
ap.add_argument("--no-capture", action="store_true")
ap.add_argument("--precision", default=None)
ap.add_argument("--ln-fp32", default=None, help="module globs (comma-separated) for the fp32 LayerNorm")
ap.add_argument("--hidden-fp32", default=None, help="module globs (comma-separated) for fp32 hidden activations")
ap.add_argument("--split", default=None, help="module globs (comma-separated) for split (hi/lo) fp32 matmuls")
ap.add_argument("--attn-fp32-acc", default=None, help="module globs for SDPA with fp32 accumulation")
ap.add_argument("--attn-matmul", default=None, help="module globs for the fp32 matmul attention")
ap.add_argument("--plan-only", action="store_true", help="skip the encoder taps and decoder evaluations")
ap.add_argument("--evals", nargs="+", type=int, default=[0, 5, 10])
ap.add_argument("--json", default=None)
args = ap.parse_args()
import torch
torch.set_num_threads(4)
from tt_diffusion_planner.device import close_device, open_device
from tt_diffusion_planner.reference.weights import find_weights_dir, load_weights
from tt_diffusion_planner.tt.model import TtDiffusionPlanner
weights = load_weights(find_weights_dir())
dev = open_device(allow_fallback=False)
g = dev.compute_with_storage_grid_size()
report = {"grid": f"{g.x}x{g.y}", "precision": args.precision, "scenes": {}}
t0 = time.perf_counter()
from tt_diffusion_planner.tt.config import globs
opts = {k: globs(v) for k, v in (("ln_fp32", args.ln_fp32), ("hidden_fp32", args.hidden_fp32),
("split", args.split), ("attn_fp32_acc", args.attn_fp32_acc),
("attn_matmul", args.attn_matmul)) if v is not None}
tt = TtDiffusionPlanner(dev, weights, debug=True, precision=args.precision, **opts)
report["options"] = tt.build.options()
report["build_s"] = round(time.perf_counter() - t0, 2)
print(f"build {report['build_s']} s, uploaded {tt.build.uploaded_bytes / 2**20:.1f} MiB", flush=True)
runner = tt.runner
from tt_diffusion_planner.tt import inputs as I
eager_cache = {}
try:
for scene in args.scenes:
with np.load(GOLDENS / f"{scene}.npz", allow_pickle=False) as z:
gold = {k: z[k] for k in z.files if k != "__meta__"}
raw = {k: gold[f"in.{k}"] for k in C.INPUT_NAMES}
prep = hp.prepare(raw, weights.normalization.observation)
res = {}
if args.plan_only:
args.evals = []
# encoder taps (eager x2)
t1 = time.perf_counter()
taps = tt.encoder_taps(prep, eager=True)
taps2 = taps if args.plan_only else tt.encoder_taps(prep, eager=True)
res["encoder_eager_s"] = round(time.perf_counter() - t1, 2)
res["encoder_deterministic"] = all(np.array_equal(taps[k], taps2[k]) for k in taps)
for name, _ in C.TOKEN_LAYOUT:
rows = np.flatnonzero(gold[f"host.valid.{name}"])
if rows.size:
res[f"enc.{name}"] = stats(taps[f"enc.{name}"][rows], gold[f"enc.{name}"][rows])
for c in ("ego", "neighbor", "lane", "route", "polygon", "line_string"):
rows = gold[f"enc.{c}.pre.rows"]
if rows.size:
for part in ("pre", "mixer"):
res[f"enc.{c}.{part}"] = stats(taps[f"enc.{c}.{part}"][rows], gold[f"enc.{c}.{part}"])
tok = np.flatnonzero(gold["host.token_valid"])
for name in ["enc.tokens"] + [f"enc.fusion.{i}" for i in range(6)] + ["enc.encoding"]:
res[name] = stats(taps[name][tok], gold[name][tok])
# decode_once, teacher forced
arows = gold["dec.rows"]
for k in args.evals:
outs = [tt_decode(tt, runner, prep, gold, k) for _ in range(2)]
res[f"dec.eval{k}.deterministic"] = bool(np.array_equal(outs[0], outs[1]))
res[f"dec.eval{k}"] = stats(outs[0][arows][:, 1:], gold["dec.out"][k][:, 1:])
# plan (eager x2)
t1 = time.perf_counter()
p1 = runner.run_eager("plan", inputs=I.plan_inputs(prep))
p2 = runner.run_eager("plan", inputs=I.plan_inputs(prep))
res["plan_eager_s"] = round(time.perf_counter() - t1, 2)
res["plan_deterministic"] = all(np.array_equal(p1[k], p2[k]) for k in p1)
eager_cache[scene] = (prep, p1)
final = p1["final_x0"].reshape(352, 324)[:321].reshape(321, 81, 4)
res["plan.final_x0"] = stats(final[arows], gold["final_x0"][arows])
res["plan.logit"] = {"dev": [round(float(v), 4) for v in p1["logit"].reshape(-1)[:5]],
"ref": [round(float(v), 4) for v in gold["turn.logit"]]}
out = hp.make_output(final, p1["logit"].reshape(-1)[:5], prep, weights.normalization,
{k: v[3] for k, v in hp.RUNTIME_PARAMS.items()})
dpos = np.hypot(*(out.poses[:, :2] - gold["out.trajectory"][:, :2]).T)
res["plan.ego_max_m"], res["plan.ego_mean_m"] = round(float(dpos.max()), 4), round(float(dpos.mean()), 4)
res["plan.turn_equal"] = int(out.turn_indicator["command"]) == int(gold["out.turn_command"])
if gold["out.predicted_agents"].shape[0]:
dxy = out.predicted_agents[..., :2] - gold["out.predicted_agents"][..., :2]
per = np.hypot(*dxy.transpose(2, 0, 1))
res["plan.nb_median_max_m"] = round(float(np.median(per.max(axis=1))), 4)
report["scenes"][scene] = res
print(json.dumps({scene: res}, indent=1), flush=True)
if not args.no_capture:
t1 = time.perf_counter()
runner.capture()
report["capture_s"] = round(time.perf_counter() - t1, 2)
report["timings_ms"] = runner.timings_ms
for scene, (prep, p1) in eager_cache.items():
rep = runner("plan", inputs=I.plan_inputs(prep))
report["scenes"][scene]["plan_replay_equals_eager"] = all(np.array_equal(rep[k], p1[k]) for k in p1)
import ttnn
ttnn.synchronize_device(dev)
t2 = time.perf_counter()
runner.replay("plan", 5)
ttnn.synchronize_device(dev)
report["scenes"][scene]["plan_replay_ms"] = round((time.perf_counter() - t2) / 5 * 1e3, 3)
print(json.dumps({k: v for k, v in report.items() if k != "scenes"}, indent=1, default=str), flush=True)
print(json.dumps({s: {k: v for k, v in r.items() if k.startswith("plan_")}
for s, r in report["scenes"].items()}, indent=1), flush=True)
finally:
tt.release()
close_device(dev)
if args.json:
Path(args.json).parent.mkdir(parents=True, exist_ok=True)
Path(args.json).write_text(json.dumps(report, indent=1, default=str) + "\n")
return 0
def tt_decode(tt, runner, prep, gold, k):
return tt.decode_once(prep, gold["dec.x_in"][k], float(gold["dec.t"][k]), encoding=gold["enc.encoding"],
eager=True)
if __name__ == "__main__":
sys.exit(main())