#!/usr/bin/env python3 # SPDX-License-Identifier: Apache-2.0 """Minimal standalone repro: ``ttnn.grid_sample`` at the METEOR lift shapes, replayed in a trace loop (no METEOR code). bin/devrun -t 900 -- python -u repro_grid_sample_eth.py --dispatch eth --iters 2000 --jsonl out.jsonl bin/devrun -t 900 -- python -u repro_grid_sample_eth.py --dispatch worker ... # A/B One trace = ``--per-trace`` x (grid_sample of x64 [8, 108, 192, 64] and of x96 [8, 108, 192, 96], bf16 ROW_MAJOR, with one fp32 grid [8, 400, 250, 2]; bilinear, zeros, align_corners=False; DRAM outputs [8, 400, 250, C]). The grid is ``--grid file.npy`` (e.g. a real METEOR lift grid: normalised coords in [-2, 2], ~half out of image) or uniform random in [-1.2, 1.2]. Each replay is followed by ``synchronize_device`` and a fsync'd JSON heartbeat; every ``--check-every`` replays the outputs are read back and hashed (must not change). faulthandler dumps the Python stack when a replay stalls for ``--stall-s``. Nothing arms tt-triage or an operation timeout. """ from __future__ import annotations import argparse import faulthandler import hashlib import json import os import sys import time import numpy as np 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=2000) ap.add_argument("--per-trace", type=int, default=1, help="grid_sample pairs per trace") ap.add_argument("--grid", help=".npy fp32 [8, 400, 250, 2] grid (default: random)") ap.add_argument("--check-every", type=int, default=100) 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() log(fh, ev="open", dispatch=a.dispatch, cqs=a.cqs, grid=f"{g.x}x{g.y}", per_trace=a.per_trace, grid_file=a.grid) rng = np.random.default_rng(0) if a.grid: grid_np = np.load(a.grid).astype(np.float32).reshape(8, 400, 250, 2) else: grid_np = rng.uniform(-1.2, 1.2, (8, 400, 250, 2)).astype(np.float32) mk = lambda arr, dt: ttnn.from_torch(torch.from_numpy(arr), dtype=dt, layout=ttnn.ROW_MAJOR_LAYOUT, # noqa: E731 device=dev, memory_config=ttnn.DRAM_MEMORY_CONFIG) x64 = mk(rng.standard_normal((8, 108, 192, 64)).astype(np.float32), ttnn.bfloat16) x96 = mk(rng.standard_normal((8, 108, 192, 96)).astype(np.float32), ttnn.bfloat16) grid = mk(grid_np, ttnn.float32) def body(): outs = [] for _ in range(a.per_trace): for x in (x64, x96): outs.append(ttnn.grid_sample(x, grid, mode="bilinear", padding_mode="zeros", align_corners=False, memory_config=ttnn.DRAM_MEMORY_CONFIG)) return outs warm = body() # compile; then free, capture ttnn.synchronize_device(dev) for t in warm: ttnn.deallocate(t) tid = ttnn.begin_trace_capture(dev, cq_id=0) outs = body() 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())