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