File size: 6,981 Bytes
c699c4c
 
 
 
 
026da6c
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
026da6c
c699c4c
026da6c
 
 
c699c4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""Phase-split benchmark of the TT_FUSED serving path (baseline profiling helper).

Run from code/ after `source tools/chipenv.sh <chip> <venv>`; PYTHONPATH must start with code/.
  SP_DISPATCH=auto|eth|worker (default auto = ETH, 1 CQ)  SP_N_ITER=200   (compute grid is always 12x10)  python bench_fused.py
Measures, each as a median over N iterations with a device sync per iteration (or a pipelined
mean where noted):
  trace_only      execute_trace + synchronize (input already resident)          -> device forward
  trace_pipelined N back-to-back execute_trace, one sync                       -> device throughput
  h2d             copy_host_to_device_tensor of the prepared 614 KB input + sync
  d2h_desc        to_torch of the resident descriptor output (RM bf16 [1,1,4800,256])
  d2h_nms         to_torch of the resident NMS map (RM bf16 [1,1,480,640])
  host_convert    reshape/permute/float of the readbacks
  forward         run_fused() end to end (prepare host input + H2D + trace + D2H + convert)
  post            postprocess_from_nms_map (threshold/border/top-k/grid_sample)
SP_EAGER_PROFILE=1: run only the eager fused graph a few times (for the device profiler; no trace).
"""
import os
import statistics
import time

import torch
import ttnn

from models.reference.superpoint_reference import get_natural_input, load_reference_model  # noqa: E402
from models.tt import postprocess as _post  # noqa: E402
from models.tt.superpoint_ttnn import TtSuperPoint  # noqa: E402


def _open():
    from models.tt.device_open import open_ttnn_device, resolve_dispatch

    mode = resolve_dispatch(os.environ.get("SP_DISPATCH", "auto"))  # auto: ETH + 1 CQ + 12x10
    kw = dict(l1_small_size=32 * 1024, trace_region_size=32 << 20)
    return open_ttnn_device(0, dispatch=mode, **kw), mode


def _med(f, n):
    ts = []
    for _ in range(n):
        t0 = time.perf_counter()
        f()
        ts.append((time.perf_counter() - t0) * 1e3)
    return statistics.median(ts), min(ts)


def main():
    n = int(os.environ.get("SP_N_ITER", "200"))
    dev, mode = _open()
    g = dev.compute_with_storage_grid_size()
    try:
        tm = load_reference_model()
        px = get_natural_input(batch_size=1)
        sp = TtSuperPoint(tm, dev, fused=True)
        tt_in = sp.allocate_input(1)
        if os.environ.get("SP_EAGER_PROFILE") == "1":
            sp.load_input(tt_in, px)  # natural frame, so the keypoint stages see real data
            for _ in range(3):
                outs = sp.build_fused_graph(tt_in, 1)
                ttnn.synchronize_device(dev)
                outs.deallocate()
            ttnn.ReadDeviceProfiler(dev)
            print("eager profile done")
            return
        sp.run_fused(tt_in, px)
        ttnn.synchronize_device(dev)
        sp.capture_trace(tt_in, 1)
        host_in = sp.prepare_host_input(px)
        sp.load_input_prepared(tt_in, host_in)
        ttnn.synchronize_device(dev)
        tid = sp.trace_id
        outs = sp._trace_outputs
        for _ in range(10):
            ttnn.execute_trace(dev, tid, cq_id=0, blocking=False)
        ttnn.synchronize_device(dev)

        def trace_once():
            ttnn.execute_trace(dev, tid, cq_id=0, blocking=False)
            ttnn.synchronize_device(dev)

        r = {}
        r["trace_only"] = _med(trace_once, n)
        t0 = time.perf_counter()
        for _ in range(n):
            ttnn.execute_trace(dev, tid, cq_id=0, blocking=False)
        ttnn.synchronize_device(dev)
        r["trace_pipelined"] = ((time.perf_counter() - t0) * 1e3 / n,) * 2

        def h2d():
            ttnn.copy_host_to_device_tensor(host_in, tt_in, cq_id=0)
            ttnn.synchronize_device(dev)

        r["h2d"] = _med(h2d, n)
        r["prepare_host_input"] = _med(lambda: sp.prepare_host_input(px), n)
        r["d2h_desc"] = _med(lambda: ttnn.to_torch(outs.d_out), n)
        r["d2h_nms"] = _med(lambda: ttnn.to_torch(outs.nms_map), n)
        r["d2h_scores_fallback"] = _med(lambda: ttnn.to_torch(outs.s_sm), max(10, n // 10))
        r["read_fused_total"] = _med(lambda: sp._read_fused(outs, 1, 480, 640, sp.nms_radius_traced), n)
        r["forward"] = _med(lambda: sp.run_fused(tt_in, px), n)
        res = sp.run_fused(tt_in, px)
        r["post"] = _med(
            lambda: _post.postprocess_from_nms_map(
                res.nms_map, res.descriptors_nchw, keypoint_threshold=0.005, max_keypoints=1024,
                border_removal_distance=4, with_descriptors=True,
            ),
            n,
        )
        # Full request (everything after JPEG decode/resize): host prep + H2D + trace + D2H + host
        # post-processing (threshold/border/top-k/descriptor sampling).
        kw = dict(keypoint_threshold=0.005, max_keypoints=1024, border_removal_distance=4)

        def e2e_legacy_api():
            r_ = sp.run_fused(tt_in, px)
            return _post.postprocess_from_nms_map(r_.nms_map, r_.descriptors_nchw, with_descriptors=True, **kw)[0]

        r["e2e_run_fused+post"] = _med(e2e_legacy_api, n)
        if hasattr(sp, "run_fused_keypoints"):
            r["e2e_keypoints"] = _med(lambda: sp.run_fused_keypoints(tt_in, sp.prepare_host_input(px), **kw), n)
            a = e2e_legacy_api()
            b = sp.run_fused_keypoints(tt_in, sp.prepare_host_input(px), **kw)
            same = torch.equal(a[0], b[0]) and torch.equal(a[1], b[1])
            print(f"e2e_keypoints vs run_fused+post: kp/scores identical={same} n_kp={a[0].shape[0]} "
                  f"desc max|diff|={float((a[2] - b[2]).abs().max()):.2e}")
        if hasattr(sp, "run_fused_keypoints_kpc") and sp._gather_tid is not None:
            r["e2e_kpc"] = _med(lambda: sp.run_fused_keypoints_kpc(tt_in, sp.prepare_host_input(px), **kw), n)
            a = e2e_legacy_api()
            c = sp.run_fused_keypoints_kpc(tt_in, sp.prepare_host_input(px), **kw)
            b = sp.run_fused_keypoints(tt_in, sp.prepare_host_input(px), **kw)
            print(f"e2e_kpc vs run_fused+post: kp/scores identical={torch.equal(a[0], c[0]) and torch.equal(a[1], c[1])} "
                  f"n_kp={c[0].shape[0]} desc max|diff|={float((a[2] - c[2]).abs().max()):.2e}; "
                  f"vs run_fused_keypoints: all identical={all(torch.equal(x, y) for x, y in zip(b, c))}")
        if hasattr(sp, "prepare_host_input_u8") and sp._gather_tid is not None:
            # the served request path: uint8 R plane in (server/app.py::_preprocess_r8)
            r8 = torch.round(px[0, 0] * 255.0).to(torch.uint8).numpy()
            r["e2e_kpc_u8"] = _med(lambda: sp.run_fused_keypoints_kpc(tt_in, sp.prepare_host_input_u8(r8), **kw), n)
        print(f"mode={mode} grid={g.x}x{g.y} n_iter={n}")
        for k, (med, mn) in r.items():
            print(f"{k:22s} median={med:8.3f} ms  min={mn:8.3f} ms")
        sp.release()
    finally:
        ttnn.close_device(dev)


if __name__ == "__main__":
    main()