superpoint-p150 / code /bench_fused.py
changh95's picture
p150 ETH-dispatch compliance (2026-10-05): default ETH dispatch, 1 CQ, 12x10 in Python API and server; numbers re-measured
026da6c verified
Raw History Blame Contribute Delete
6.98 kB
# 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()