bevformer-p150 / examples /quickstart.py
changh95's picture
tt-model push bevformer-p150 (container)
ddaad99 verified
Raw History Blame Contribute Delete
4.32 kB
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Quickstart: the model card's Python snippet on the shipped sample, on one Blackhole p150.
pip install -e . # once, from the repo root, on top of an environment that has ttnn (tt-metal)
python examples/quickstart.py [sample.json] [--out-dir examples/output]
Writes <out-dir>/quickstart.json (the same JSON as POST /predict) and <out-dir>/quickstart_bev.png (a bird's-eye view
of the detections in the network's LIDAR_TOP frame). The input is a sample manifest (six camera images, a
calibration, the stream with the ego pose: ``code/tt_bevformer/samples/*.json``); the default, the shipped synthetic
test pattern, is found relative to this file (runs from any directory); an input given on the command line is
relative to the current directory.
"""
import argparse
import json
import math
from pathlib import Path
REPO = Path(__file__).resolve().parents[1]
ap = argparse.ArgumentParser()
ap.add_argument("input", nargs="?", default=str(REPO / "code" / "tt_bevformer" / "samples" / "synthetic_6cam.json"))
ap.add_argument("--out-dir", default=str(REPO / "examples" / "output"))
ap.add_argument("--device-id", type=int, default=0)
args = ap.parse_args()
out_dir = Path(args.out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
# --- the model card snippet --------------------------------------------------------------------------------------
from tt_bevformer import BEVFormer, load_sample
with BEVFormer.from_pretrained(device_id=args.device_id) as model: # weights -> HF cache, traces captured
# six cameras [F, FR, FL, B, BL, BR] + calibration; stream={...} carries the ego pose for the BEV history
out = model(**load_sample(args.input))
for d in out.to_dicts():
print(f'{d["label"]:10s} {out.meta["class_names"][d["label_id"]]:13s} {d["score"]:.3f} centre {d["center"]} '
f'size {d["size"]} yaw {d["yaw"]:+.2f} v {d["velocity"]}')
# ------------------------------------------------------------------------------------------------------------------
(out_dir / "quickstart.json").write_text(json.dumps(out.to_dict(), indent=1))
def bev_png(dets, class_names, path, rng=40.0, size=640):
"""Bird's-eye view of the detections around the ego vehicle in LIDAR_TOP (x right, y forward = up), 10 m rings.
Boxes are in the node's convention: size [w, l, h], the footprint w along the yaw axis, l across it."""
from PIL import Image, ImageDraw
vehicle, ped, cyc = (57, 135, 229), (217, 89, 38), (25, 158, 112)
colours = {"car": vehicle, "truck": vehicle, "bus": vehicle, "trailer": vehicle, "construction_vehicle": vehicle,
"pedestrian": ped, "bicycle": cyc, "motorcycle": cyc}
im = Image.new("RGB", (size, size), (26, 26, 25))
d = ImageDraw.Draw(im, "RGBA")
s = size / (2 * rng)
def px(x, y):
return ((rng + x) * s, (rng - y) * s)
for r in range(10, int(rng) + 1, 10):
(u0, v0), (u1, v1) = px(-r, r), px(r, -r)
d.ellipse([u0, v0, u1, v1], outline=(44, 44, 42), width=1)
for det in dets:
x, y, _ = det["center"]
w, l, _ = det["size"]
c, sn = math.cos(det["yaw"]), math.sin(det["yaw"])
q = [px(x + c * a * w - sn * b * l, y + sn * a * w + c * b * l)
for a, b in ((0.5, 0.5), (0.5, -0.5), (-0.5, -0.5), (-0.5, 0.5))]
name = class_names[det["label_id"]]
col = colours.get(name, (195, 194, 183))
d.polygon(q, fill=col + (70,), outline=col)
d.line([px(x, y), px(x - sn * l / 2, y + c * l / 2)], fill=col, width=2) # heading (the length axis)
d.text((max(q[0][0], q[1][0]) + 3, min(q[0][1], q[1][1]) - 12), f'{name} {det["score"]:.2f}',
fill=(195, 194, 183))
ego = [px(a * 0.9, b * 2.05 - 0.9) for a, b in ((1, 1), (1, -1), (-1, -1), (-1, 1))]
d.polygon(ego, fill=(255, 255, 255))
d.text((6, size - 16), f"{2 * rng:.0f} m x {2 * rng:.0f} m, rings every 10 m, up = forward (LIDAR_TOP +y)",
fill=(137, 135, 129))
im.save(path)
bev_png(out.to_dicts(), out.meta["class_names"], out_dir / "quickstart_bev.png")
print(f"{len(out)} detections -> {out_dir / 'quickstart.json'}, {out_dir / 'quickstart_bev.png'} "
f"timing_ms={ {k: round(v, 2) for k, v in out.timing_ms.items()} }")