File size: 9,228 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 164 165 166 167 168 | # 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())
|