File size: 5,230 Bytes
51defdc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
#!/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())