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()
|