# SPDX-License-Identifier: Apache-2.0 """Precision experiment for the MLP-Mixer blocks on the device (development tool; run under ``bin/devrun``). bin/devrun -t 1200 -- python code/scripts/precision_exp.py [--json logs/diffusion-planner/precision_exp.json] Part A: ``ttnn.layer_norm`` vs :func:`tt.layers.layer_norm_fp32` on the real offset-dominated mixer rows (the reference ``enc.neighbor.pre`` of the samples, block 0 ``norm1``) against float64. Part B: the 6 mixer blocks of ego / neighbour / lane run alone on the reference's own block input (``enc..pre``, teacher forcing) under several option sets, compared with the reference ``enc..mixer`` on the valid entities. Eager (no trace). """ from __future__ import annotations import argparse import json import sys from pathlib import Path import numpy as np HERE = Path(__file__).resolve() sys.path.insert(0, str(HERE.parents[1])) 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" CONFIGS = { "default": dict(), "ln_fp32": dict(ln_fp32=("enc.mixer.*",)), "hidden_fp32": dict(hidden_fp32=("enc.mixer.*",)), "ln_hidden_fp32": dict(ln_fp32=("enc.mixer.*",), hidden_fp32=("enc.mixer.*",)), "all_fp32_w": dict(ln_fp32=("enc.mixer.*",), hidden_fp32=("enc.mixer.*",), precision="enc.mixer.*=HiFi4+fp32:w=fp32:a=fp32"), "stream_bf16": dict(precision="enc.mixer.*=HiFi4+fp32:w=bf16:a=bf16"), } def st(dev, ref): s = error_stats(np.asarray(dev, np.float64), np.asarray(ref, np.float64)) return {"pcc": round(s["pcc"], 7), "rel_l2": float(f"{s['rel_l2']:.4g}"), "max_abs": float(f"{s['max_abs']:.4g}")} def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--scenes", nargs="+", default=["kashiwanoha_dense", "straight_road"]) ap.add_argument("--cats", nargs="+", default=["ego", "neighbor", "lane"]) ap.add_argument("--configs", nargs="*", default=list(CONFIGS)) ap.add_argument("--json", default=None) args = ap.parse_args() import torch import ttnn 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.encoder import ENTITIES, MixerTrunk from tt_diffusion_planner.tt.layers import Build, layer_norm_fp32, policy from tt_diffusion_planner.ttaw.precision import compute_kernel_config from tt_diffusion_planner.ttaw.tensors import to_device, to_numpy p = load_weights(find_weights_dir()).params gold = {} for s in args.scenes: with np.load(GOLDENS / f"{s}.npz", allow_pickle=False) as z: gold[s] = {k: z[k] for k in z.files if k.startswith(("enc.", "host.valid"))} dev = open_device(allow_fallback=False) report = {"ln": {}, "mixer": {}} try: # ---- A: LayerNorm accuracy on real rows g = p["encoder.neighbor_encoder.blocks.0.norm1.gamma"] b = p["encoder.neighbor_encoder.blocks.0.norm1.beta"] gd, bd = (to_device(v.reshape(1, 1, 1, -1), dev, "float32") for v in (g, b)) for s in args.scenes: x = gold[s]["enc.neighbor.pre"].reshape(1, 1, -1, C.MIXER_CHANNELS).astype(np.float32) xx = x.astype(np.float64) ref = (xx - xx.mean(-1, keepdims=True)) / np.sqrt(xx.var(-1, keepdims=True) + 1e-5) * g + b tx = to_device(x, dev, "float32") fused = to_numpy(ttnn.layer_norm(tx, epsilon=1e-5, weight=gd, bias=bd, compute_kernel_config=compute_kernel_config("HiFi4", fp32_acc=True))) dec = to_numpy(layer_norm_fp32(tx, gd, bd)) report["ln"][s] = {"fused_fp32": st(fused, ref), "decomposed_fp32": st(dec, ref), "row_offset_ratio": float(np.abs(xx.mean(-1)).mean() / xx.std(-1).mean())} print("LN", s, json.dumps(report["ln"][s]), flush=True) # ---- B: mixer blocks alone, teacher-forced for name in args.configs: cfg = dict(CONFIGS[name]) build = Build(dev, policy(spec=cfg.pop("precision", None)), **cfg) for cat in args.cats: trunk = MixerTrunk(build, p, cat) for s in args.scenes: rows = gold[s][f"enc.{cat}.pre.rows"] if rows.size == 0: continue x0 = np.zeros((1, ENTITIES[cat], C.MIXER_TOKENS, C.MIXER_CHANNELS), np.float32) x0[0, rows] = gold[s][f"enc.{cat}.pre"] out = to_numpy(trunk.mix(to_device(x0, dev, trunk.stream)))[0, rows] r = st(out, gold[s][f"enc.{cat}.mixer"]) report["mixer"].setdefault(name, {}).setdefault(cat, {})[s] = r print(f"MIX {name:15s} {cat:9s} {s:18s} {json.dumps(r)}", flush=True) del trunk # ---- C: the island on the device (TF32-like vs split matmuls) -> pre error, and the float64 CPU mixer on # the device pre (separates the input error from the device mixer) from tt_diffusion_planner.host import pipeline as hp from tt_diffusion_planner.reference.model import Encoder, torch_params from tt_diffusion_planner.tt import inputs as I ref_enc = Encoder(torch_params(p, torch.float64)) weights = load_weights(find_weights_dir()) def cpu_mix(x, cat): N = f"encoder.{cat}_encoder" x = torch.from_numpy(np.asarray(x, np.float64)) for i in range(C.MIXER_DEPTH): B = f"{N}.blocks.{i}" x = x + ref_enc.mlp(ref_enc.ln(x, f"{B}.norm1").transpose(1, 2), f"{B}.tokens_mlp").transpose(1, 2) x = x + ref_enc.mlp(ref_enc.ln(x, f"{B}.norm2"), f"{B}.channels_mlp") return x.numpy() preps = {} for s in args.scenes: with np.load(GOLDENS / f"{s}.npz", allow_pickle=False) as z: raw = {k: z[f"in.{k}"] for k in C.INPUT_NAMES} preps[s] = I.plan_inputs(hp.prepare(raw, weights.normalization.observation)) for mode, split in (("tf32", ()), ("split", ("enc.island.*",))): build = Build(dev, policy(), ln_fp32=("enc.mixer.*",), split=split) for cat in ("ego", "neighbor"): trunk = MixerTrunk(build, p, cat) for s in args.scenes: rows = gold[s][f"enc.{cat}.pre.rows"] pre_t = trunk.pre(to_device(preps[s][f"{cat}_x"], dev, "float32")) pre = to_numpy(pre_t)[0, rows] mix = to_numpy(trunk.mix(pre_t))[0, rows] r = {"pre": st(pre, gold[s][f"enc.{cat}.pre"]), "mix_device": st(mix, gold[s][f"enc.{cat}.mixer"]), "mix_cpu64_on_device_pre": st(cpu_mix(pre, cat), gold[s][f"enc.{cat}.mixer"])} report.setdefault("island", {}).setdefault(mode, {}).setdefault(cat, {})[s] = r print(f"ISL {mode:6s} {cat:9s} {s:18s} {json.dumps(r)}", flush=True) del trunk finally: close_device(dev) if args.json: Path(args.json).write_text(json.dumps(report, indent=1) + "\n") return 0 if __name__ == "__main__": sys.exit(main())