Download code/scripts/repro_grid_sample_eth.py from changh95/meteor-p150: direct link, hf CLI and curl.
- Browser
- Download file 5.23 kB
-
https://huggingface.co/changh95/meteor-p150/resolve/main/code/scripts/repro_grid_sample_eth.py
- Command line
-
hf download hf://changh95/meteor-p150/code/scripts/repro_grid_sample_eth.py
-
curl -L -o repro_grid_sample_eth.py https://huggingface.co/changh95/meteor-p150/resolve/main/code/scripts/repro_grid_sample_eth.py
5.23 kB
| #!/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()) | |