meteor-p150 / code /scripts /precision_emulation.py
changh95's picture
tt-model push meteor-p150 (container)
51defdc verified
Raw History Blame Contribute Delete
9.4 kB
# SPDX-License-Identifier: Apache-2.0
"""CPU emulation of the device's activation precision on the class decisions (PORT_LOG.md E20; no device).
cd bundles/meteor-p150 && OMP_NUM_THREADS=4 PYTHONPATH=$PWD/code \\
/home/ubuntu/experiments/tt-models/tools/research-venv/bin/python code/scripts/precision_emulation.py image bb tt Bt
... precision_emulation.py bev
... precision_emulation.py det meteor_valday_f040 fff bff fbf ffb bbb nbb bFb nFF
... precision_emulation.py absent ff ft fT TT tt bb
``image``: the CPU reference's image branch with every conv / resize output rounded to bf16 (``b``), TF32 (``t``: 10
explicit mantissa bits, what the device's fp32 conv path computes: its unpacker rounds fp32 activations to TF32) or
kept fp32 (``f``), per (encoder, heads) pair; ``B`` = the stem output alone in bf16 (``ttnn.max_pool2d`` is bf16
only). Prints the depth / seg2d argmax agreement with the fp32 reference (PLAN 2.13 gate: >= 0.99) and the depth
top-2 logit margins. ``bev``: the lane decoder + seg refiner from the golden raw BEV with bf16 activations (the convs
whose outputs the device keeps fp32 stay fp32) -> lane agreement. ``det FRAME RDB ...`` (PORT_LOG.md E22): the det
stem + merged heads (``D``) and the box refiner (``B``) from the golden raw BEV (``R``) of FRAME, each letter one of
``f`` (fp32), ``b`` (bf16 conv / resize outputs), ``t`` (TF32), ``F`` (fp32 convs = three bf16 terms, TF32 resizes:
the device's terms3 path), and for ``R`` also ``n`` (bf16 after 0.3 % relative noise: the chained device raw) -> the
strict det3d recall / precision of the decoded boxes vs the golden outputs (``tests/box_agreement.py``) and the hm
error. ``absent CR ...`` (PORT_LOG.md E23 / I7): the image branch on an absent camera (the all-zero normalised input
of ``input_norm=imagenet``) with conv outputs ``C`` and resize outputs ``R`` rounded as above, plus ``T`` (TF32 by
TRUNCATION) -> seg2d / depth agreement with the variant golden's absent camera and with the device's argmax maps
(``logs/meteor/variants_device_absent_argmax.npz``, written by ``tests/test_variants_device.py``).
Goldens: ``research/meteor/goldens``.
Measured (PandaSet 019 f40): image bb 0.9819 / 0.9908, bt 0.9856 / 0.9948, tb 0.9913 / 0.9917, tt 0.9981 / 0.9988,
Bt 0.9874 / 0.9951; bev lane 0.99966 (bf16) / 0.99996 (fp32). det (valday #40, strict recall / precision): fff
1.0 / 1.0, bff 1.0 / 1.0, fbf 0.971 / 1.0, ffb 1.0 / 1.0, bbb 0.914 / 1.0, nbb 0.914 / 1.0, bFb 1.0 / 1.0, nFF
1.0 / 1.0 (nuScenes nbb 0.821 / 0.885, nFF 1.0 / 1.0): the bf16 det stem is the loss. absent (seg2d / depth vs the
reference; vs the device): ff 1.0 / 1.0; fT 0.9975 / 0.9993; tt 0.9966 / 0.9981; TT 0.9856 / 0.9901 (0.9927 / 0.9933
vs the device); bb 0.9724 / 0.9800; the device 0.9832 / 0.9913 (camera 6): it behaves like TF32-truncated conv outputs.
"""
from __future__ import annotations
import os
import sys
from pathlib import Path
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
GOLD = Path(os.environ.get("METEOR_GOLDENS", "/home/ubuntu/experiments/tt-models/research/meteor/goldens"))
FP32_OUT_BEV = {"dec/out/out.3", "lane_branch/lane_branch.6", "refiner/seg/out/out.3"}
def rnd(t, m):
import torch
if m == "f":
return t
if m == "T": # TF32 by truncation
a = t.contiguous().numpy().view(np.uint32) & np.uint32(0xFFFFE000)
return torch.from_numpy(a.view(np.float32).copy())
if m == "b":
return t.to(torch.bfloat16).to(torch.float32)
a = t.contiguous().numpy().view(np.uint32).astype(np.uint64)
a = ((a + 0x1000) & 0xFFFFE000).astype(np.uint32) # round to 10 explicit mantissa bits (TF32)
return torch.from_numpy(a.view(np.float32))
def main() -> None:
import torch
from tt_meteor.reference.model import MeteorNet
from tt_meteor.reference.weights import MeteorWeights
torch.set_num_threads(int(os.environ.get("OMP_NUM_THREADS", "4")))
what = sys.argv[1] if len(sys.argv) > 1 else "image"
net = MeteorNet(MeteorWeights(), check_graph=False)
z = np.load(GOLD / "pandaset_019_f40" / "taps.npz")
cur = {"m": "b", "stem": None}
conv0, resize0 = net.conv, net.resize
def conv(x, module, relu=False):
y = conv0(x, module, relu)
if what == "bev" and module in FP32_OUT_BEV:
return y
if module == "stem/stem.0" and cur["stem"]:
return rnd(y, cur["stem"])
return rnd(y, cur["m"])
if what == "det":
return det(net, sys.argv[2], sys.argv[3:])
if what == "absent":
return absent(net, sys.argv[2:])
net.conv = conv
net.resize = lambda x, hw: rnd(resize0(x, hw), cur["m"])
with torch.inference_mode():
if what == "bev":
raw = rnd(torch.from_numpy(z["bev.raw"].astype(np.float32)), "b")
for m in ("b", "f"):
cur["m"] = m
ll = net.seg_refiner(net.lane_decoder(raw))
lane = torch.argmax(ll, dim=1).numpy().astype(np.uint8)
print(f"bev {m}: lane agreement {(lane == z['lane']).mean():.5f}")
return
x = net.normalize(torch.from_numpy(z["input.imgs"]))
for cfg in sys.argv[2:] or ["bb", "tt"]:
enc, head = cfg[0], cfg[1]
cur["stem"] = "b" if enc == "B" else None
cur["m"] = "t" if enc == "B" else enc
f = net.image_encoder(x)
cur["m"], cur["stem"] = head, None
dl, _ = net.depth_head(f)
seg = net.seg2d_head(f)
d = torch.argmax(dl.reshape(1, 8, 64, 108, 192), dim=2).numpy()
s = torch.argmax(torch.clamp(seg, -30, 30).reshape(1, 8, 21, 108, 192), dim=2).numpy()
print(f"encoder {enc} heads {head}: depth {(d == z['depth']).mean():.4f} seg2d {(s == z['seg2d']).mean():.4f}",
flush=True)
srt = np.sort(z["depth.logits"].astype(np.float32), axis=1)
margin = srt[:, -1] - srt[:, -2]
print("depth top-2 margin quantiles (0.05, 0.1, 0.25, 0.5):",
np.round(np.quantile(margin, [0.05, 0.1, 0.25, 0.5]), 4).tolist())
def det(net, frame: str, configs) -> None:
"""``det`` mode (module docstring)."""
import torch
from tt_meteor.host.postprocess import PostConfig
from tt_meteor.tests.box_agreement import box_agreement, decode_3d
g = np.load(GOLD / frame / "taps.npz")
conv0, resize0 = net.conv, net.resize
cur = {"m": "f"}
conv_mode = lambda m: "f" if m == "F" else m # noqa: E731
resize_mode = lambda m: "t" if m == "F" else m # noqa: E731
net.conv = lambda x, module, relu=False: rnd(conv0(x, module, relu), conv_mode(cur["m"]))
net.resize = lambda x, hw: rnd(resize0(x, hw), resize_mode(cur["m"]))
cfg = PostConfig()
ref = decode_3d({"hm": g["hm"], "reg": g["reg"]}, cfg)
torch.manual_seed(0)
for c in configs or ["bbb", "bFF"]:
r_m, d_m, b_m = c
with torch.inference_mode():
raw = torch.from_numpy(g["bev.raw"].astype(np.float32))
raw = rnd(raw + torch.randn_like(raw) * raw.abs() * 0.003, "b") if r_m == "n" else rnd(raw, r_m)
cur["m"] = d_m
_, hm_pre, reg_pre = net.det_stem(raw)
cur["m"] = b_m
hm, reg = net.box_refiner(hm_pre, reg_pre)
r = box_agreement(decode_3d({"hm": hm.numpy(), "reg": reg.numpy()}, cfg), ref, cfg)
err = np.abs(hm.numpy() - g["hm"]).ravel()
print(f"{frame} raw {r_m} det {d_m} refiner {b_m}: recall {r['recall']:.4f} precision {r['precision']:.4f} "
f"hm max abs {err.max():.4f}", flush=True)
def absent(net, configs) -> None:
"""``absent`` mode (module docstring)."""
import torch
g = np.load(GOLD / "pandaset_019_f40_imagenet_linear" / "taps.npz")
cam = int(np.flatnonzero(~np.asarray(g["input.present"], bool))[0])
ref_s, ref_d = g["seg2d"][0][cam], g["depth"][0][cam]
dev_path = Path(os.environ.get("TT_MODELS_ROOT", "/home/ubuntu/experiments/tt-models")) / "logs" / "meteor" / \
"variants_device_absent_argmax.npz"
dev = None
if dev_path.is_file():
with np.load(dev_path) as d:
i = int(np.flatnonzero(d["absent"] == cam)[0])
dev = {"seg2d": d["seg2d"][i], "depth": d["depth"][i]}
conv0, resize0 = net.conv, net.resize
cur = {"c": "f", "r": "f"}
net.conv = lambda x, module, relu=False: rnd(conv0(x, module, relu), cur["c"])
net.resize = lambda x, hw: rnd(resize0(x, hw), cur["r"])
x = torch.zeros(1, 3, *g["input.imgs"].shape[-2:]) # imagenet: absent = zero AFTER normalising
for c in configs or ["ff", "TT"]:
cur["c"], cur["r"] = c[0], c[1]
with torch.inference_mode():
f = net.image_encoder(x)
dl, _ = net.depth_head(f)
seg = net.seg2d_head(f)
s = torch.argmax(torch.clamp(seg, -30, 30), dim=1).numpy()[0]
d = torch.argmax(dl.reshape(1, -1, *s.shape), dim=1).numpy()[0]
msg = f"absent cam {cam} conv {c[0]} resize {c[1]}: seg2d {(s == ref_s).mean():.4f} depth {(d == ref_d).mean():.4f}"
if dev is not None:
msg += f" | vs device seg2d {(s == dev['seg2d']).mean():.4f} depth {(d == dev['depth']).mean():.4f}"
print(msg, flush=True)
if __name__ == "__main__":
main()