#!/usr/bin/env python3 # SPDX-License-Identifier: Apache-2.0 """Minimal standalone repro (plain ttnn, no METEOR code): the separable bilinear x2 resize of the METEOR lift, captured in one trace and replayed in a tight loop. On a Blackhole p150 with ETH dispatch the replay intermittently never completes (logs/meteor/hangfix/ETH_DISPATCH_HANG_ISSUE.md). bin/devrun -t 600 -- python -u repro_resize2d_eth.py --dispatch eth --iters 3000 --jsonl out.jsonl bin/devrun -t 600 -- python -u repro_resize2d_eth.py --dispatch worker ... # A/B Op chain (= ttaw.ops.upsample.Resize2d, in (400, 250) -> out (800, 500), C = 96, N = 1, fp32 operands, HiFi4 + fp32 accumulation), on an input x [1, 1, 100000, 96] bf16 TILE in DRAM: RM -> reshape [1, 400, 250, 96] -> TILE -> typecast fp32 -> transpose(-2, -1) [1, 400, 96, 250] -> matmul A_w^T [250, 500] -> transpose(-2, -1) [1, 400, 500, 96] -> RM -> reshape [1, 1, 400, 48000] -> TILE -> matmul A_h [800, 400] @ [1, 1, 400, 48000] -> [1, 1, 800, 48000] fp32 -> typecast bf16 -> RM -> reshape [1, 1, 400000, 96] -> TILE ``--part w`` stops after the W pass, ``--part h`` runs only the H-pass matmul (on a persistent [1, 1, 400, 48000] fp32 TILE input). Each replay is followed by ``synchronize_device`` and a fsync'd JSON heartbeat; outputs are hashed every ``--check-every`` replays; faulthandler dumps the stack when a replay stalls. No tt-triage, no op timeout. """ from __future__ import annotations import argparse import faulthandler import hashlib import json import os import sys import time import numpy as np H, W, H2, W2 = 400, 250, 800, 500 def interp(n_in: int, n_out: int) -> np.ndarray: """1-D bilinear, align_corners=False (half-pixel, clamped at 0), as F.interpolate.""" s = n_in / n_out src = np.maximum((np.arange(n_out) + 0.5) * s - 0.5, 0.0) i0 = np.minimum(np.floor(src).astype(np.int64), n_in - 1) i1 = np.minimum(i0 + 1, n_in - 1) lam = src - i0 a = np.zeros((n_out, n_in)) np.add.at(a, (np.arange(n_out), i0), 1.0 - lam) np.add.at(a, (np.arange(n_out), i1), lam) return a.astype(np.float32) def log(fh, **rec): rec["t"] = round(time.time(), 3) line = json.dumps(rec) print(line, flush=True) if fh: fh.write(line + "\n") fh.flush() os.fsync(fh.fileno()) def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--dispatch", choices=["eth", "worker"], default="eth") ap.add_argument("--cqs", type=int, default=1) ap.add_argument("--iters", type=int, default=3000) ap.add_argument("--per-trace", type=int, default=1, help="resize chains per trace") ap.add_argument("--part", choices=["full", "w", "h", "w_rm", "wh", "tail", "mid", "mid_h", "h_tail", "t_typecast", "t_untilize", "t_reshape", "t_tilize", "t_reshape_tile"], default="full", help="w: up to the W-pass transpose; w_rm: + RM/reshape/TILE; wh: + H matmul (no tail); " "tail: typecast/RM/reshape/TILE of a persistent z; mid: RM/reshape/TILE of a persistent " "[1,400,500,C] y; mid_h: mid + H matmul; h_tail: H matmul + tail; t_*: ONE tail op on a " "persistent input (typecast fp32->bf16 TILE [800, 500C]; untilize bf16 [800, 500C]; RM " "reshape [800, 500C] -> [400000, C]; tilize RM [400000, C])") ap.add_argument("--channels", type=int, default=96, help="C (the H-pass matmul N is 500 * C; 96 -> 48000)") ap.add_argument("--dtype", choices=["float32", "bfloat16"], default="float32", help="matmul operand dtype") ap.add_argument("--fidelity", choices=["HiFi4", "HiFi3", "HiFi2", "LoFi"], default="HiFi4") ap.add_argument("--no-fp32-acc", action="store_true", help="fp32_dest_acc_en=False") ap.add_argument("--dest-cols", type=int, default=0, help="t_reshape: destination last dim (default C: [800, 500C] -> [400000, C]); dest page = 2 B x it") ap.add_argument("--reshape", choices=["rm", "tile", "tail_tile"], default="rm", help="rm: TILE -> RM -> reshape -> TILE (as Resize2d); tile: ttnn.reshape on the TILE tensor") ap.add_argument("--check-every", type=int, default=200) ap.add_argument("--stall-s", type=float, default=60.0) ap.add_argument("--jsonl") a = ap.parse_args() faulthandler.enable() fh = open(a.jsonl, "a") if a.jsonl else None faulthandler.dump_traceback_later(600, repeat=True) import torch import ttnn core = ttnn.DispatchCoreType.ETH if a.dispatch == "eth" else ttnn.DispatchCoreType.WORKER dev = ttnn.open_device(device_id=0, dispatch_core_config=ttnn.DispatchCoreConfig(core), num_command_queues=a.cqs, l1_small_size=32768, trace_region_size=64 << 20) rc = 0 try: g = dev.compute_with_storage_grid_size() C = a.channels dt = ttnn.float32 if a.dtype == "float32" else ttnn.bfloat16 log(fh, ev="open", dispatch=a.dispatch, cqs=a.cqs, grid=f"{g.x}x{g.y}", part=a.part, per_trace=a.per_trace, channels=C, n=W2 * C, dtype=a.dtype, fidelity=a.fidelity, fp32_acc=not a.no_fp32_acc, reshape=a.reshape, dest_cols=a.dest_cols or None) cfg = ttnn.init_device_compute_kernel_config(dev.arch(), math_fidelity=getattr(ttnn.MathFidelity, a.fidelity), fp32_dest_acc_en=not a.no_fp32_acc, packer_l1_acc=False, math_approx_mode=False) to_dev = lambda arr, dt, lay: ttnn.from_torch(torch.from_numpy(np.ascontiguousarray(arr)), dtype=dt, # noqa layout=lay, device=dev, memory_config=ttnn.DRAM_MEMORY_CONFIG) rng = np.random.default_rng(0) aw_t = to_dev(interp(W, W2).T, dt, ttnn.TILE_LAYOUT) # [250, 500] ah = to_dev(interp(H, H2), dt, ttnn.TILE_LAYOUT) # [800, 400] x_in = to_dev(rng.standard_normal((1, 1, H * W, C)).astype(np.float32), ttnn.bfloat16, ttnn.TILE_LAYOUT) y_in = to_dev(rng.standard_normal((1, 1, H, W2 * C)).astype(np.float32), dt, ttnn.TILE_LAYOUT) z_in = to_dev(rng.standard_normal((1, 1, H2, W2 * C)).astype(np.float32), dt, ttnn.TILE_LAYOUT) \ if a.part == "tail" else None yw_in = to_dev(rng.standard_normal((1, H, W2, C)).astype(np.float32), dt, ttnn.TILE_LAYOUT) \ if a.part in ("mid", "mid_h") else None one = None if a.part.startswith("t_"): z32 = rng.standard_normal((1, 1, H2, W2 * C)).astype(np.float32) one = {"t_typecast": lambda: to_dev(z32, ttnn.float32, ttnn.TILE_LAYOUT), "t_untilize": lambda: to_dev(z32, ttnn.bfloat16, ttnn.TILE_LAYOUT), "t_reshape": lambda: to_dev(z32, ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT), "t_tilize": lambda: to_dev(z32.reshape(1, 1, H2 * W2, C), ttnn.bfloat16, ttnn.ROW_MAJOR_LAYOUT), "t_reshape_tile": lambda: to_dev(z32, ttnn.bfloat16, ttnn.TILE_LAYOUT), }[a.part]() def tail(z): if z.dtype != ttnn.bfloat16: z = ttnn.typecast(z, ttnn.bfloat16) if a.reshape in ("tile", "tail_tile"): return ttnn.reshape(z, (1, 1, H2 * W2, C)) z = ttnn.to_layout(z, ttnn.ROW_MAJOR_LAYOUT) z = ttnn.reshape(z, (1, 1, H2 * W2, C)) return ttnn.to_layout(z, ttnn.TILE_LAYOUT) def mid(y): if a.reshape == "tile": return ttnn.reshape(y, (1, 1, H, W2 * C)) y = ttnn.to_layout(y, ttnn.ROW_MAJOR_LAYOUT) y = ttnn.reshape(y, (1, 1, H, W2 * C)) return ttnn.to_layout(y, ttnn.TILE_LAYOUT) def chain(): if a.part == "t_typecast": return ttnn.typecast(one, ttnn.bfloat16) if a.part == "t_untilize": return ttnn.to_layout(one, ttnn.ROW_MAJOR_LAYOUT) if a.part == "t_reshape": d = a.dest_cols or C return ttnn.reshape(one, (1, 1, H2 * W2 * C // d, d)) if a.part == "t_reshape_tile": return ttnn.reshape(one, (1, 1, H2 * W2, C)) if a.part == "t_tilize": return ttnn.to_layout(one, ttnn.TILE_LAYOUT) if a.part == "h": return ttnn.matmul(ah, y_in, compute_kernel_config=cfg) if a.part == "h_tail": return tail(ttnn.matmul(ah, y_in, compute_kernel_config=cfg)) if a.part == "tail": return tail(z_in) if a.part == "mid": return mid(yw_in) if a.part == "mid_h": return ttnn.matmul(ah, mid(yw_in), compute_kernel_config=cfg) x = ttnn.to_layout(x_in, ttnn.ROW_MAJOR_LAYOUT) x = ttnn.reshape(x, (1, H, W, C)) x = ttnn.to_layout(x, ttnn.TILE_LAYOUT) if dt != ttnn.bfloat16: x = ttnn.typecast(x, dt) xt = ttnn.transpose(x, -2, -1) y = ttnn.matmul(xt, aw_t, compute_kernel_config=cfg) y = ttnn.transpose(y, -2, -1) if a.part == "w": return y y = mid(y) if a.part == "w_rm": return y z = ttnn.matmul(ah, y, compute_kernel_config=cfg) if a.part == "wh": return z return tail(z) warm = [chain() for _ in range(a.per_trace)] # compile, then capture ttnn.synchronize_device(dev) del warm tid = ttnn.begin_trace_capture(dev, cq_id=0) outs = [chain() for _ in range(a.per_trace)] ttnn.end_trace_capture(dev, tid, cq_id=0) ttnn.synchronize_device(dev) log(fh, ev="captured") ref = None for i in range(a.iters): faulthandler.dump_traceback_later(a.stall_s, repeat=True) t0 = time.perf_counter() ttnn.execute_trace(dev, tid, cq_id=0, blocking=False) ttnn.synchronize_device(dev) rec = {"ev": "iter", "i": i, "ms": round((time.perf_counter() - t0) * 1e3, 2)} if i % a.check_every == 0 or i == a.iters - 1: h = hashlib.sha256() for t in outs: h.update(ttnn.to_torch(t).float().numpy().tobytes()) d = h.hexdigest()[:16] rec["digest"] = d ref = ref or d if d != ref: rec["MISMATCH"] = ref rc = 1 log(fh, **rec) faulthandler.cancel_dump_traceback_later() log(fh, ev="done", iters=a.iters, rc=rc) ttnn.release_trace(dev, tid) finally: faulthandler.dump_traceback_later(300, repeat=True) ttnn.close_device(dev) faulthandler.cancel_dump_traceback_later() log(fh, ev="closed") return rc if __name__ == "__main__": sys.exit(main())