Download code/models/tests/test_superpoint.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 43.5 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tests/test_superpoint.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/tests/test_superpoint.py
-
curl -L -o test_superpoint.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tests/test_superpoint.py
43.5 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """SuperPoint benchmark + accuracy test. | |
| Run from ``code/`` with the tt-metal tree's python_env (the ``device`` / ``device_params`` | |
| fixtures and ``--device-id`` come from ``code/conftest.py``; tt-metal's own conftest cannot be | |
| loaded next to this repo because ``code/models`` shadows its namespace ``models`` package): | |
| <tree>/python_env/bin/python -m pytest -s -q --device-id=0 \ | |
| models/tests/test_superpoint.py::test_superpoint_benchmark # legacy path (fused=False; run_benchmark.sh) | |
| <tree>/python_env/bin/python -m pytest -s -q --device-id=0 \ | |
| models/tests/test_superpoint.py::test_superpoint_fused # fused path (default) A/B via TT_FUSED_STAGES | |
| Prints: | |
| inference_speed=<fps> | |
| accuracy=<percent_of_baseline_PCC> | |
| peak_dram=<bytes> | |
| """ | |
| from __future__ import annotations | |
| import os | |
| import time | |
| import pytest | |
| import torch | |
| import torch.nn.functional as F | |
| import ttnn | |
| from loguru import logger | |
| from models.reference.superpoint_reference import ( | |
| load_reference_model, | |
| get_dummy_input, | |
| get_natural_input, | |
| ) | |
| from models.tt.superpoint_ttnn import ( | |
| TtSuperPoint, | |
| TtConv2D, | |
| KEYPOINT_DIM, | |
| DESCRIPTOR_DIM, | |
| ) | |
| TRACE_REGION_SIZE = 6 * 1024 * 1024 # 6 MB trace region | |
| def _pcc(a: torch.Tensor, b: torch.Tensor) -> float: | |
| a = a.detach().float().flatten() | |
| b = b.detach().float().flatten() | |
| if a.numel() == 0 or b.numel() == 0: | |
| return 1.0 | |
| a = a - a.mean() | |
| b = b - b.mean() | |
| denom = (a.norm() * b.norm()).item() | |
| if denom == 0: | |
| return 1.0 | |
| return float((a @ b).item() / denom) | |
| def _topk_keypoints(score_map: torch.Tensor, k: int) -> torch.Tensor: | |
| """Return the (y, x) coordinates of the k highest-scoring pixels.""" | |
| flat = score_map.flatten() | |
| k_eff = min(k, flat.numel()) | |
| _, idx = torch.topk(flat, k_eff) | |
| h = score_map.shape[-2] | |
| w = score_map.shape[-1] | |
| return torch.stack([idx // w, idx % w], dim=1) | |
| def _keypoint_set_metrics(tt_scores: torch.Tensor, ref_scores: torch.Tensor, k: int = 500, tol: int = 2): | |
| """Compare top-K keypoints with a pixel tolerance. | |
| tt_scores, ref_scores: (H, W) post-NMS dense score maps. | |
| tol: match if tt keypoint is within ``tol`` pixels of a reference keypoint. | |
| Returns (recall, precision, f1). | |
| """ | |
| tt_kp = _topk_keypoints(tt_scores, k) | |
| ref_kp = _topk_keypoints(ref_scores, k) | |
| if tt_kp.numel() == 0 or ref_kp.numel() == 0: | |
| return 0.0, 0.0, 0.0 | |
| # For each ref keypoint, is there any tt keypoint within tol pixels? | |
| d = torch.cdist(ref_kp.float(), tt_kp.float(), p=torch.inf) | |
| ref_matched = (d.min(dim=1).values <= tol).float().mean().item() | |
| tt_matched = (d.min(dim=0).values <= tol).float().mean().item() | |
| recall = ref_matched | |
| precision = tt_matched | |
| f1 = 2 * recall * precision / max(recall + precision, 1e-9) | |
| return recall, precision, f1 | |
| def _device_to_host_post(tt_model, s_sm, d_norm, b, h, w): | |
| """Convert device outputs (softmax already applied on device) to NCHW host tensors.""" | |
| enc_h, enc_w = h // 8, w // 8 | |
| scores_nhwc = ttnn.to_torch(s_sm).reshape(b, enc_h, enc_w, KEYPOINT_DIM) | |
| descriptors_nhwc = ttnn.to_torch(d_norm).reshape(b, enc_h, enc_w, DESCRIPTOR_DIM) | |
| scores_nchw = scores_nhwc.permute(0, 3, 1, 2).contiguous().float() | |
| descriptors_nchw = descriptors_nhwc.permute(0, 3, 1, 2).contiguous().float() | |
| return scores_nchw, descriptors_nchw | |
| def _device_to_host_post_with_nms(tt_model, s_pooled, d_norm, b, h, w): | |
| """Hot-loop D2H: pull the device-NMS'd map (single-channel, row-major) | |
| and descriptors. ``sp_eq_mul_mask`` + on-device channel-0 slice make the | |
| D2H payload 32× smaller than the 32-padded eq-mul output. | |
| """ | |
| enc_h, enc_w = h // 8, w // 8 | |
| descriptors_nhwc = ttnn.to_torch(d_norm).reshape(b, enc_h, enc_w, DESCRIPTOR_DIM) | |
| nms_scores = ttnn.to_torch(s_pooled).reshape(b, h, w).float() | |
| descriptors_nchw = descriptors_nhwc.permute(0, 3, 1, 2).contiguous().float() | |
| return nms_scores, descriptors_nchw | |
| def test_superpoint_benchmark(device, height, width, input_kind): | |
| torch.manual_seed(0) | |
| torch_model = load_reference_model() | |
| if input_kind == "natural": | |
| pixel_values = get_natural_input(batch_size=1, height=height, width=width) | |
| else: | |
| pixel_values = get_dummy_input(batch_size=1, height=height, width=width) | |
| with torch.no_grad(): | |
| _ = torch_model(pixel_values=pixel_values) # keep weights loaded on CPU | |
| # This is the LEGACY benchmark (run_benchmark.sh, results.tsv history): pin the knob off | |
| # explicitly -- TT_FUSED defaults to the fused path since 2026-09-13. | |
| tt_model = TtSuperPoint(torch_model, device, input_height=height, input_width=width, fused=False) | |
| b = 1 | |
| # Persistent device input tensor (filled via copy_host_to_device_tensor). | |
| tt_in = tt_model.allocate_input(batch_size=b) | |
| # A.1 device fold+NMS exists but is opt-in — the ~6 ms device fold | |
| # overhead doesn't pay off in Python-composed form vs the 24-36 ms | |
| # host single-pass NMS. A fused C++ kernel would flip this. | |
| trace_nms = os.environ.get("SP_TRACE_NMS", "0") == "1" | |
| # Propagate the flag so run_device_compute (which reads the env var at | |
| # call time) returns the 3-tuple (s, s_pooled, d_norm) we expect here. | |
| os.environ["SP_TRACE_NMS"] = "1" if trace_nms else "0" | |
| def _do_warmup(): | |
| out = tt_model.run_device_compute(tt_in, b=b) | |
| if trace_nms: | |
| sw, pw, dw = out | |
| ttnn.synchronize_device(device) | |
| ttnn.deallocate(sw); ttnn.deallocate(pw); ttnn.deallocate(dw) | |
| else: | |
| sw, dw = out | |
| ttnn.synchronize_device(device) | |
| ttnn.deallocate(sw); ttnn.deallocate(dw) | |
| # Warmup/compile: first full forward compiles the graph. | |
| t0 = time.perf_counter() | |
| tt_model.load_input(tt_in, pixel_values) | |
| _do_warmup() | |
| t_compile = time.perf_counter() - t0 | |
| logger.info(f"compile/warmup time: {t_compile:.3f}s (SP_TRACE_NMS={int(trace_nms)})") | |
| use_trace = os.environ.get("SP_NO_TRACE", "0") != "1" | |
| if use_trace: | |
| # Capture trace of the device compute graph. | |
| tt_model.load_input(tt_in, pixel_values) | |
| tid = ttnn.begin_trace_capture(device, cq_id=0) | |
| out = tt_model.run_device_compute(tt_in, b=b) | |
| if trace_nms: | |
| s, s_pooled, d_norm = out | |
| else: | |
| s, d_norm = out | |
| s_pooled = None | |
| ttnn.end_trace_capture(device, tid, cq_id=0) | |
| # Warmup the trace execution once (allocator setup). | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=True) | |
| # Produce one forward result via the traced path for PCC comparison. | |
| tt_model.load_input(tt_in, pixel_values) | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=True) | |
| tt_scores_nchw, tt_desc_nchw = _device_to_host_post(tt_model, s, d_norm, b, height, width) | |
| # Pre-build the host bf16 tensor once. ttnn.from_torch with a bf16 cast | |
| # costs ~10 ms/iter if repeated in the hot loop — moving that out of | |
| # the loop lets the per-iter H2D become pure PCIe DMA. | |
| host_input = tt_model.prepare_host_input(pixel_values) | |
| n_iter = int(os.environ.get("SP_N_ITER", "10")) | |
| # Pure-compute upper bound: input already resident on device, timed | |
| # loop is just traced replay. | |
| tt_model.load_input_prepared(tt_in, host_input) | |
| ttnn.synchronize_device(device) | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=False) | |
| ttnn.synchronize_device(device) | |
| fps_compute_only = n_iter / (time.perf_counter() - t0) | |
| # Timed iterations — traced replay, then the next frame's H2D on the same queue (CQ0; in order, | |
| # so the upload never overtakes the trace that reads the input). | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=False) | |
| tt_model.load_input_prepared(tt_in, host_input, cq_id=0) | |
| ttnn.synchronize_device(device) | |
| _ = _device_to_host_post(tt_model, s, d_norm, b, height, width) | |
| elapsed = time.perf_counter() - t0 | |
| fps = n_iter / elapsed | |
| # Second timed loop: end-to-end throughput on one command queue. | |
| # CQ0 runs the current trace, then the next frame's H2D (same queue, | |
| # in order); D2H + host post-processing for the current frame run | |
| # afterward on the Python thread. Per-iter cost: H2D + compute + D2H + post. | |
| e2e_phase_times = {"h2d": 0.0, "compute": 0.0, "d2h": 0.0, "post": 0.0} | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| tp0 = time.perf_counter() | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=False) | |
| tt_model.load_input_prepared(tt_in, host_input, cq_id=0) | |
| ttnn.synchronize_device(device) | |
| tp2 = time.perf_counter() | |
| if trace_nms: | |
| nms_scores, desc_host = _device_to_host_post_with_nms( | |
| tt_model, s_pooled, d_norm, b, height, width | |
| ) | |
| tp3 = time.perf_counter() | |
| for i in range(b): | |
| kp, sc = tt_model._extract_keypoints_single(nms_scores[i : i + 1]) | |
| if kp.shape[0] > 0: | |
| _ = tt_model._sample_descriptors(kp[None], desc_host[i : i + 1], scale=8) | |
| tp4 = time.perf_counter() | |
| else: | |
| scores_host, desc_host = _device_to_host_post(tt_model, s, d_norm, b, height, width) | |
| tp3 = time.perf_counter() | |
| scores_full = tt_model._decode_keypoints(scores_host, apply_nms=True) | |
| for i in range(b): | |
| kp, sc = tt_model._extract_keypoints_single(scores_full[i : i + 1]) | |
| if kp.shape[0] > 0: | |
| _ = tt_model._sample_descriptors(kp[None], desc_host[i : i + 1], scale=8) | |
| tp4 = time.perf_counter() | |
| e2e_phase_times["compute"] += tp2 - tp0 # trace + next H2D on CQ0 | |
| e2e_phase_times["d2h"] += tp3 - tp2 | |
| e2e_phase_times["post"] += tp4 - tp3 | |
| elapsed_e2e = time.perf_counter() - t0 | |
| fps_e2e = n_iter / elapsed_e2e | |
| # Paper-matching slice: forward + D2H + descriptor sampling only (no NMS). | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| tt_model.load_input_prepared(tt_in, host_input) | |
| ttnn.execute_trace(device, tid, cq_id=0, blocking=True) | |
| scores_host, desc_host = _device_to_host_post(tt_model, s, d_norm, b, height, width) | |
| scores_pre = tt_model._decode_keypoints(scores_host, apply_nms=False) | |
| for i in range(b): | |
| flat = scores_pre[i].flatten() | |
| _, idx = torch.topk(flat, 1000) | |
| w_ = scores_pre.shape[-1] | |
| kp = torch.stack([idx // w_, idx % w_], dim=1).flip(1).to(torch.float32) | |
| _ = tt_model._sample_descriptors(kp[None], desc_host[i : i + 1], scale=8) | |
| elapsed_match = time.perf_counter() - t0 | |
| fps_match_paper = n_iter / elapsed_match | |
| else: | |
| # Fallback for profilers: no trace, so per-op markers are visible. | |
| # Force SP_TRACE_NMS=0 here to keep the 2-tuple return shape. | |
| os.environ["SP_TRACE_NMS"] = "0" | |
| tt_model.load_input(tt_in, pixel_values) | |
| s, d_norm = tt_model.run_device_compute(tt_in, b=b) | |
| ttnn.synchronize_device(device) | |
| tt_scores_nchw, tt_desc_nchw = _device_to_host_post(tt_model, s, d_norm, b, height, width) | |
| ttnn.deallocate(s) | |
| ttnn.deallocate(d_norm) | |
| tid = None | |
| n_iter = int(os.environ.get("SP_N_ITER", "10")) | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| tt_model.load_input(tt_in, pixel_values) | |
| s, d_norm = tt_model.run_device_compute(tt_in, b=b) | |
| ttnn.deallocate(s) | |
| ttnn.deallocate(d_norm) | |
| ttnn.synchronize_device(device) | |
| elapsed = time.perf_counter() - t0 | |
| fps = n_iter / elapsed | |
| fps_e2e = fps # no-trace path doesn't separately time e2e | |
| fps_match_paper = fps | |
| fps_compute_only = fps | |
| # Build full SuperPoint output structure for accuracy comparison. | |
| tt_scores_pre_nms = tt_model._decode_keypoints(tt_scores_nchw, apply_nms=False) | |
| with torch.no_grad(): | |
| enc = torch_model.encoder(torch_model.extract_one_channel_pixel_values(pixel_values))[0] | |
| ks = torch_model.keypoint_decoder.relu(torch_model.keypoint_decoder.conv_score_a(enc)) | |
| ks = torch_model.keypoint_decoder.conv_score_b(ks) | |
| ks = F.softmax(ks, 1)[:, :-1] | |
| _, _, h_, w_ = ks.shape | |
| ks = ks.permute(0, 2, 3, 1).reshape(1, h_, w_, 8, 8) | |
| ref_score_pre = ks.permute(0, 1, 3, 2, 4).reshape(1, h_ * 8, w_ * 8) | |
| ref_desc_full = F.normalize( | |
| torch_model.descriptor_decoder.conv_descriptor_b( | |
| torch_model.descriptor_decoder.relu(torch_model.descriptor_decoder.conv_descriptor_a(enc)) | |
| ), | |
| p=2, | |
| dim=1, | |
| ) | |
| score_pcc = _pcc(tt_scores_pre_nms, ref_score_pre) | |
| desc_pcc = _pcc(tt_desc_nchw, ref_desc_full) | |
| accuracy = min(score_pcc, desc_pcc) * 100.0 | |
| # Keypoint-set overlap after NMS — the real downstream metric for | |
| # SuperPoint consumers (matching, SLAM, etc.). | |
| tt_scores_nms = tt_model._simple_nms(tt_scores_pre_nms, tt_model.nms_radius)[0] | |
| with torch.no_grad(): | |
| ref_scores_nms = tt_model._simple_nms(ref_score_pre, tt_model.nms_radius)[0] | |
| recall_500, precision_500, f1_500 = _keypoint_set_metrics(tt_scores_nms, ref_scores_nms, k=500, tol=2) | |
| print(f"input_kind={input_kind}") | |
| print(f"inference_speed={fps:.4f} fps") | |
| print(f"inference_speed_compute_only={fps_compute_only:.4f} fps") | |
| print(f"inference_speed_e2e={fps_e2e:.4f} fps") | |
| if use_trace: | |
| for k, v in e2e_phase_times.items(): | |
| print(f"e2e_phase_ms_{k}={v / n_iter * 1000:.3f}") | |
| print(f"inference_speed_match_paper={fps_match_paper:.4f} fps") | |
| print(f"accuracy={accuracy:.4f}") | |
| print(f"score_pcc={score_pcc:.6f}") | |
| print(f"descriptor_pcc={desc_pcc:.6f}") | |
| print(f"keypoint_recall@500_tol2={recall_500:.4f}") | |
| print(f"keypoint_precision@500_tol2={precision_500:.4f}") | |
| print(f"keypoint_f1@500_tol2={f1_500:.4f}") | |
| print(f"peak_dram={0}") | |
| assert torch.isfinite(tt_scores_pre_nms).all() | |
| assert torch.isfinite(tt_desc_nchw).all() | |
| if tid is not None: | |
| ttnn.release_trace(device, tid) | |
| # --------------------------------------------------------------------------- TT_FUSED path | |
| # Device test for the hardware pass of the opt/superpoint-p150-megakernel branch (never run on | |
| # the host; see DEVICE_VALIDATION.md). Same accuracy gates as the benchmark above, plus the | |
| # bit-identity gates the fused reformulations promise: | |
| # * eager fused graph == traced replay (descriptors and NMS map, torch.equal) | |
| # * device NMS-T map == host fold_scores(s_sm, r) + simple_nms on the SAME traced s_sm (torch.equal) | |
| # * score PCC >= 0.997, descriptor PCC >= 0.999, keypoint F1 >= 0.9879 @ top-500 / 2 px (natural image; | |
| # the legacy path measures F1 0.98796 = recall 0.9820 / precision 0.9940, "98.80%" on the card) | |
| # A/B a single stage with TT_FUSED_STAGES (e.g. "" = trace-only, "wide", "wide,nms", ...). | |
| FUSED_TRACE_REGION_SIZE = 32 * 1024 * 1024 | |
| def test_superpoint_fused(device, height, width, input_kind): | |
| from models.tt import postprocess as _post | |
| torch.manual_seed(0) | |
| torch_model = load_reference_model() | |
| if input_kind == "natural": | |
| pixel_values = get_natural_input(batch_size=1, height=height, width=width) | |
| else: | |
| pixel_values = get_dummy_input(batch_size=1, height=height, width=width) | |
| from models.tt import fused_host as _fh | |
| if "u8" in _fh.fused_stages(): | |
| # the u8 input stage takes 8-bit images (the served input domain): a random 8-bit | |
| # image / 255 instead of uniform floats (round 1, 2026-10-03) | |
| pixel_values = torch.round(pixel_values * 255.0) / 255.0 | |
| stages = os.environ.get("TT_FUSED_STAGES") # None -> all stages | |
| tt_model = TtSuperPoint(torch_model, device, input_height=height, input_width=width, fused=True) | |
| logger.info(f"TT_FUSED stages: {sorted(tt_model.fused_stages)} (TT_FUSED_STAGES={stages!r})") | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| # 1) eager compile pass (JIT + conv weight preparation), 2) capture, 3) traced replay. | |
| t0 = time.perf_counter() | |
| eager = tt_model.run_fused(tt_in, pixel_values) | |
| ttnn.synchronize_device(device) | |
| t_compile = time.perf_counter() - t0 | |
| t0 = time.perf_counter() | |
| tt_model.capture_trace(tt_in, b=1) | |
| t_capture = time.perf_counter() - t0 | |
| traced = tt_model.run_fused(tt_in, pixel_values) | |
| ttnn.synchronize_device(device) | |
| logger.info(f"compile {t_compile:.2f}s, capture {t_capture*1e3:.1f} ms, trace_id={tt_model.trace_id}") | |
| assert tt_model.trace_id is not None | |
| # Same graph, same input -> eager and traced outputs must be identical. | |
| assert torch.equal(eager.descriptors_nchw, traced.descriptors_nchw), "eager vs traced descriptors differ" | |
| if traced.nms_map is not None: | |
| assert eager.nms_map is not None and torch.equal(eager.nms_map, traced.nms_map), "eager vs traced NMS map differ" | |
| # Scores via the fallback readback (any radius != traced -> scores_nchw from the traced s_sm). | |
| fallback = tt_model.run_fused(tt_in, pixel_values, nms_radius=tt_model.nms_radius_traced + 1) | |
| assert fallback.scores_nchw is not None and fallback.nms_map is None | |
| tt_scores_nchw, tt_desc_nchw = fallback.scores_nchw, fallback.descriptors_nchw | |
| tt_scores_pre_nms = _post.fold_scores(tt_scores_nchw, None) | |
| # Device NMS-T must be bit-identical to the host fold + simple_nms of the SAME s_sm. | |
| if traced.nms_map is not None: | |
| host_nms = _post.fold_scores(tt_scores_nchw, tt_model.nms_radius_traced) | |
| n_diff = int((traced.nms_map != host_nms).sum()) | |
| print(f"fused_nms_map_mismatches={n_diff}") | |
| assert n_diff == 0, f"device NMS-T map differs from host simple_nms in {n_diff} pixels" | |
| tt_scores_nms = traced.nms_map[0] | |
| else: | |
| tt_scores_nms = _post.fold_scores(tt_scores_nchw, tt_model.nms_radius)[0] | |
| # Replay determinism over a few iterations + timing (H2D + execute_trace + D2H + host convert). | |
| n_iter = int(os.environ.get("SP_N_ITER", "20")) | |
| host_input = tt_model.prepare_host_input(pixel_values) | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| tt_model.load_input_prepared(tt_in, host_input) | |
| ttnn.execute_trace(device, tt_model.trace_id, cq_id=0, blocking=False) | |
| ttnn.synchronize_device(device) | |
| compute_ms = (time.perf_counter() - t0) / n_iter * 1e3 # H2D + trace, no D2H | |
| t0 = time.perf_counter() | |
| for _ in range(n_iter): | |
| again = tt_model.run_fused(tt_in, pixel_values) | |
| forward_ms = (time.perf_counter() - t0) / n_iter * 1e3 | |
| assert torch.equal(again.descriptors_nchw, traced.descriptors_nchw) | |
| if traced.nms_map is not None: | |
| assert torch.equal(again.nms_map, traced.nms_map) | |
| # Host post-processing time on the fused result (what the server does after device_forward). | |
| t0 = time.perf_counter() | |
| if traced.nms_map is not None: | |
| kp, sc, desc = _post.postprocess_from_nms_map( | |
| traced.nms_map, traced.descriptors_nchw, keypoint_threshold=0.005, max_keypoints=1024, | |
| border_removal_distance=4, with_descriptors=True, | |
| )[0] | |
| else: | |
| kp, sc, desc = _post.postprocess_keypoints( | |
| tt_scores_nchw, tt_desc_nchw, nms_radius=tt_model.nms_radius, keypoint_threshold=0.005, | |
| max_keypoints=1024, border_removal_distance=4, with_descriptors=True, | |
| )[0] | |
| post_ms = (time.perf_counter() - t0) * 1e3 | |
| # On-device post-processing fast path (kpc stage): same keypoints/scores as the host path, and | |
| # bit-identical descriptors to the NHWC host sampler (the legacy grid_sample differs by fp32 | |
| # rounding of the bilinear sum only). Several max_keypoints values exercise top-k and buckets. | |
| if getattr(tt_model, "_gather_tid", None) is not None: | |
| hin = tt_model.prepare_host_input(pixel_values) | |
| for mk in (1024, 100, -1, 0): | |
| kw = dict(keypoint_threshold=0.005, max_keypoints=mk, border_removal_distance=4) | |
| a = _post.postprocess_from_nms_map(traced.nms_map, traced.descriptors_nchw, with_descriptors=True, **kw)[0] | |
| b = tt_model.run_fused_keypoints(tt_in, hin, **kw) | |
| c = tt_model.run_fused_keypoints_kpc(tt_in, hin, **kw) | |
| assert torch.equal(a[0], c[0]) and torch.equal(a[1], c[1]), f"kpc keypoints differ (max_keypoints={mk})" | |
| assert all(torch.equal(x, y) for x, y in zip(b, c)), f"kpc differs from run_fused_keypoints ({mk})" | |
| dd = float((a[2] - c[2]).abs().max()) if c[0].shape[0] else 0.0 | |
| assert dd < 1e-5, dd | |
| print(f"fused_kpc_identical[max_kp={mk}]=True n_kp={c[0].shape[0]} desc_maxdiff_vs_grid_sample={dd:.2e}") | |
| # Reference (identical to test_superpoint_benchmark). | |
| with torch.no_grad(): | |
| enc = torch_model.encoder(torch_model.extract_one_channel_pixel_values(pixel_values))[0] | |
| ks = torch_model.keypoint_decoder.relu(torch_model.keypoint_decoder.conv_score_a(enc)) | |
| ks = torch_model.keypoint_decoder.conv_score_b(ks) | |
| ks = F.softmax(ks, 1)[:, :-1] | |
| _, _, h_, w_ = ks.shape | |
| ks = ks.permute(0, 2, 3, 1).reshape(1, h_, w_, 8, 8) | |
| ref_score_pre = ks.permute(0, 1, 3, 2, 4).reshape(1, h_ * 8, w_ * 8) | |
| ref_desc_full = F.normalize( | |
| torch_model.descriptor_decoder.conv_descriptor_b( | |
| torch_model.descriptor_decoder.relu(torch_model.descriptor_decoder.conv_descriptor_a(enc)) | |
| ), | |
| p=2, | |
| dim=1, | |
| ) | |
| ref_scores_nms = _post.simple_nms(ref_score_pre, tt_model.nms_radius)[0] | |
| score_pcc = _pcc(tt_scores_pre_nms, ref_score_pre) | |
| desc_pcc = _pcc(tt_desc_nchw, ref_desc_full) | |
| recall_500, precision_500, f1_500 = _keypoint_set_metrics(tt_scores_nms, ref_scores_nms, k=500, tol=2) | |
| desc_norm_dev = float((tt_desc_nchw.norm(dim=1) - 1.0).abs().max()) | |
| print(f"input_kind={input_kind}") | |
| print(f"fused_stages={','.join(sorted(tt_model.fused_stages)) or 'trace-only'}") | |
| print(f"fused_compile_s={t_compile:.3f}") | |
| print(f"fused_trace_capture_ms={t_capture*1e3:.2f}") | |
| print(f"fused_h2d_plus_trace_ms={compute_ms:.3f}") | |
| print(f"fused_forward_ms={forward_ms:.3f}") | |
| print(f"fused_postprocess_ms={post_ms:.3f}") | |
| print(f"fused_num_keypoints={int(kp.shape[0])}") | |
| print(f"score_pcc={score_pcc:.6f}") | |
| print(f"descriptor_pcc={desc_pcc:.6f}") | |
| print(f"descriptor_norm_max_dev={desc_norm_dev:.5f}") | |
| print(f"keypoint_recall@500_tol2={recall_500:.4f}") | |
| print(f"keypoint_precision@500_tol2={precision_500:.4f}") | |
| print(f"keypoint_f1@500_tol2={f1_500:.4f}") | |
| assert torch.isfinite(tt_scores_pre_nms).all() and torch.isfinite(tt_desc_nchw).all() | |
| if input_kind == "natural": | |
| assert score_pcc >= 0.997, score_pcc | |
| assert desc_pcc >= 0.999, desc_pcc | |
| # Legacy path on this frame (run_benchmark.sh, 2026-09-13 p150a): recall 0.9820, precision | |
| # 0.9940 -> F1 0.98796; the card's "98.80%" is that value rounded. Gate on the measured | |
| # legacy value, not on the rounded card number (0.988 would fail the legacy path too). | |
| assert f1_500 >= 0.9879, f1_500 | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_kpc_paths(device, threshold): | |
| """Every branch of the on-device post-processing (``run_fused_keypoints_kpc``) against the host | |
| post-processing of the same traced outputs: <= 1024 candidates (in-trace list + sampling, with | |
| and without top-k), > 1024 candidates (host top-k + the second sampling trace), and > 1024 | |
| keypoints kept (resident-map fallback). Keypoints/scores must be identical, descriptors | |
| bit-identical to the NHWC host sampler.""" | |
| from models.tt import postprocess as _post | |
| torch_model = load_reference_model() | |
| torch_model.config.keypoint_threshold = threshold # traced into the candidate kernel | |
| pixel_values = get_natural_input(batch_size=1, height=480, width=640) | |
| tt_model = TtSuperPoint(torch_model, device, fused=True) | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| tt_model.run_fused(tt_in, pixel_values) | |
| tt_model.capture_trace(tt_in, b=1) | |
| traced = tt_model.run_fused(tt_in, pixel_values) | |
| hin = tt_model.prepare_host_input(pixel_values) | |
| hdr_total = None | |
| for mk in (1024, 300, -1): | |
| kw = dict(keypoint_threshold=threshold, max_keypoints=mk, border_removal_distance=4) | |
| a = _post.postprocess_from_nms_map(traced.nms_map, traced.descriptors_nchw, with_descriptors=True, **kw)[0] | |
| b = tt_model.run_fused_keypoints(tt_in, hin, **kw) | |
| c = tt_model.run_fused_keypoints_kpc(tt_in, hin, **kw) | |
| if hdr_total is None: | |
| h = tt_model._kpc_last_hdr | |
| hdr_total = f"{int(h[0])} overflow={int(h[1])}" | |
| assert torch.equal(a[0], c[0]) and torch.equal(a[1], c[1]), f"keypoints differ (thr={threshold}, max_kp={mk})" | |
| assert all(torch.equal(x, y) for x, y in zip(b, c)), f"kpc != run_fused_keypoints (thr={threshold}, max_kp={mk})" | |
| dd = float((a[2] - c[2]).abs().max()) if c[0].shape[0] else 0.0 | |
| assert dd < 1e-5, dd | |
| print(f"kpc_paths thr={threshold} candidates={hdr_total} max_kp={mk} n_kp={c[0].shape[0]} identical=True desc_maxdiff_vs_grid_sample={dd:.2e}") | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_nms_radius_variants(device): | |
| """Per-request nms_radius != traced radius: the precompiled per-radius device NMS + keypoint | |
| trace must give exactly the host result (fold + simple_nms(r) + extraction on the same traced | |
| scores); descriptors equal to grid_sample up to fp32 rounding.""" | |
| from models.tt import postprocess as _post | |
| torch_model = load_reference_model() | |
| pixel_values = get_natural_input(batch_size=1, height=480, width=640) | |
| tt_model = TtSuperPoint(torch_model, device, fused=True) | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| tt_model.run_fused(tt_in, pixel_values) | |
| tt_model.capture_trace(tt_in, b=1) | |
| hin = tt_model.prepare_host_input(pixel_values) | |
| fb = tt_model.run_fused(tt_in, pixel_values, nms_radius=99) # host fallback readback (scores) | |
| for r in [int(v) for v in os.environ.get("SP_TEST_RADII", "2,3,5,8,4").split(",")]: | |
| for mk in (1024, 200): | |
| kw = dict(keypoint_threshold=0.005, max_keypoints=mk, border_removal_distance=4) | |
| ref = _post.postprocess_keypoints(fb.scores_nchw, fb.descriptors_nchw, nms_radius=r, with_descriptors=True, **kw)[0] | |
| got = tt_model.run_fused_keypoints_kpc(tt_in, hin, nms_radius=r, **kw) | |
| oa = torch.argsort(ref[1], descending=True, stable=True) | |
| ob = torch.argsort(got[1], descending=True, stable=True) | |
| assert ref[0].shape == got[0].shape, (r, mk, ref[0].shape, got[0].shape) | |
| assert torch.equal(ref[0][oa], got[0][ob]) and torch.equal(ref[1][oa], got[1][ob]), (r, mk) | |
| dd = float((ref[2][oa] - got[2][ob]).abs().max()) if got[0].shape[0] else 0.0 | |
| assert dd < 1e-5, (r, mk, dd) | |
| # the slower host-extraction variant (threshold != traced) on the same device map | |
| kw2 = dict(kw, keypoint_threshold=0.01) | |
| ref2 = _post.postprocess_keypoints(fb.scores_nchw, fb.descriptors_nchw, nms_radius=r, with_descriptors=True, **kw2)[0] | |
| got2 = tt_model.run_fused_keypoints_kpc(tt_in, hin, nms_radius=r, **kw2) | |
| o2a = torch.argsort(ref2[1], descending=True, stable=True) | |
| o2b = torch.argsort(got2[1], descending=True, stable=True) | |
| assert torch.equal(ref2[0][o2a], got2[0][o2b]) and torch.equal(ref2[1][o2a], got2[1][o2b]), (r, mk, "thr") | |
| print(f"nms_radius_variant r={r} max_kp={mk} n_kp={got[0].shape[0]} identical=True desc_maxdiff={dd:.2e} " | |
| f"(thr 0.01: n_kp={got2[0].shape[0]})") | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_kpc_runtime_params(device): | |
| """keypoint_threshold / border per request on ONE captured trace (class C parameters: a 64-byte | |
| parameter tensor the candidate kernel reads): every combination must match the host | |
| post-processing of the same traced maps exactly, including going back to an earlier value.""" | |
| from models.tt import postprocess as _post | |
| torch_model = load_reference_model() | |
| pixel_values = get_natural_input(batch_size=1, height=480, width=640) | |
| tt_model = TtSuperPoint(torch_model, device, fused=True) | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| tt_model.run_fused(tt_in, pixel_values) | |
| tt_model.capture_trace(tt_in, b=1) | |
| traced = tt_model.run_fused(tt_in, pixel_values) | |
| hin = tt_model.prepare_host_input(pixel_values) | |
| combos = [(0.005, 4), (0.02, 4), (1e-5, 4), (0.005, 0), (0.005, 8), (0.005, 60), (0.1, 2), (0.0, 4), (0.005, 4)] | |
| for r in (4, 3): | |
| for thr, border in combos: | |
| kw = dict(keypoint_threshold=thr, max_keypoints=1024, border_removal_distance=border) | |
| if r == tt_model.nms_radius_traced: | |
| a = _post.postprocess_from_nms_map(traced.nms_map, traced.descriptors_nchw, with_descriptors=True, **kw)[0] | |
| else: | |
| fb = tt_model.run_fused(tt_in, pixel_values, nms_radius=99) | |
| a = _post.postprocess_keypoints(fb.scores_nchw, fb.descriptors_nchw, nms_radius=r, with_descriptors=True, **kw)[0] | |
| c = tt_model.run_fused_keypoints_kpc(tt_in, hin, nms_radius=r, **kw) | |
| oa = torch.argsort(a[1], descending=True, stable=True) | |
| oc = torch.argsort(c[1], descending=True, stable=True) | |
| assert a[0].shape == c[0].shape, (r, thr, border, a[0].shape, c[0].shape) | |
| assert torch.equal(a[0][oa], c[0][oc]) and torch.equal(a[1][oa], c[1][oc]), (r, thr, border) | |
| dd = float((a[2][oa] - c[2][oc]).abs().max()) if c[0].shape[0] else 0.0 | |
| assert dd < 1e-5, (r, thr, border, dd) | |
| print(f"kpc_runtime_params r={r} thr={thr} border={border} n_kp={c[0].shape[0]} identical=True desc_maxdiff={dd:.2e}") | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_device_resize(device): | |
| """``rsz`` stage: the device bilinear resize (kernels/sp_resize/resize_r8.cpp) of a full-size | |
| uint8 plane writes exactly Pillow's 640x480 resize into the network input, for several source | |
| sizes (down / up / one-axis / odd, in any order, re-using variants), and the whole kpc request | |
| on it is identical to the request on the host-resized plane.""" | |
| import numpy as np | |
| from PIL import Image | |
| from models.tt.superpoint_ttnn import SourcePlane | |
| torch_model = load_reference_model() | |
| tt_model = TtSuperPoint(torch_model, device, fused=True) | |
| assert tt_model.device_resize | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| pixel_values = get_natural_input(batch_size=1, height=480, width=640) | |
| tt_model.run_fused(tt_in, pixel_values) | |
| tt_model.capture_trace(tt_in, b=1) | |
| rng = np.random.default_rng(3) | |
| nat = (pixel_values[0, 0].numpy() * 255.0).round().astype(np.uint8) | |
| kw = dict(keypoint_threshold=0.005, max_keypoints=1024, border_removal_distance=4) | |
| for (w, h) in [(1600, 900), (1920, 1080), (641, 481), (320, 240), (2000, 480), (1001, 777), (1600, 900), (3840, 2160)]: | |
| if (w, h) == (1600, 900): # a natural-looking frame (upscaled natural input) for the keypoint check | |
| src = np.asarray(Image.fromarray(nat).resize((w, h), Image.BICUBIC)) | |
| else: | |
| src = rng.integers(0, 256, (h, w), dtype=np.uint8) | |
| ref = np.asarray(Image.fromarray(src).resize((640, 480), Image.BILINEAR)) | |
| hin = tt_model.prepare_source(src) | |
| assert isinstance(hin, SourcePlane), (w, h) | |
| tt_model.load_input_prepared(tt_in, hin) | |
| got = ttnn.to_torch(tt_in).numpy().reshape(480, 640) | |
| assert np.array_equal(got, ref), (w, h, int((got != ref).sum())) | |
| a = tt_model.run_fused_keypoints_kpc(tt_in, tt_model.prepare_host_input_u8(ref), **kw) | |
| c = tt_model.run_fused_keypoints_kpc(tt_in, hin, **kw) | |
| assert all(torch.equal(x, y) for x, y in zip(a, c)), (w, h) | |
| print(f"device_resize {w}x{h}: input bit-identical to Pillow, kpc request identical (n_kp={c[0].shape[0]})") | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_cell0_matches_ttnn_conv(device, monkeypatch): | |
| """Block-0 conv_b + 2x2 pool as the custom cell-tile kernels (ConvCell0 + PoolCell0, SP_CELL0=1) | |
| against ttnn.conv2d + the round-1 pool kernel on the same conv_a output (natural frame): the | |
| bias is added the way ttnn's conv does it, so all but a handful of the 4.9M pooled values are | |
| bit-identical (the rest differ by one bf16 ulp).""" | |
| from models.tt.conv_cell import ConvCell0, PoolCell0 | |
| from models.tt.pool_kernels import U8ToBf16 | |
| monkeypatch.setenv("SP_CELL0", "0") | |
| # the ttnn block-0 conv needs the L1 that the descriptor head's resident weights would take | |
| monkeypatch.setenv("SP_DESC_HEAD", "0") | |
| torch_model = load_reference_model() | |
| tt_model = TtSuperPoint(torch_model, device, fused=True) | |
| tt_in = tt_model.allocate_input(batch_size=1) | |
| pixel_values = get_natural_input(batch_size=1, height=480, width=640) | |
| ttnn.copy_host_to_device_tensor(tt_model.prepare_host_input(pixel_values), tt_in) | |
| u8 = U8ToBf16(device) | |
| conv_a, conv_b, _ = tt_model.l1_convs[0] | |
| x = u8(tt_in, tt_model._cell_input_memory_config(1), tt_model._cell_input_shape(1)) | |
| xa, _, _ = conv_a(x, 480, 80, 1) | |
| ttnn.deallocate(x) | |
| yr, _, _ = conv_b(ttnn.experimental.view(xa, [1, 1, 307200, 64]), 480, 640, 1) | |
| pr = tt_model._pool2x2(yr, 480, 640) | |
| ref = ttnn.to_torch(pr).reshape(-1, 64).float() | |
| ttnn.deallocate(pr) | |
| ttnn.deallocate(yr) | |
| conv_a.conv_config.output_layout = ttnn.TILE_LAYOUT | |
| x = u8(tt_in, tt_model._cell_input_memory_config(1), tt_model._cell_input_shape(1)) | |
| xt, _, _ = conv_a(x, 480, 80, 1) | |
| ttnn.deallocate(x) | |
| blk = torch_model.encoder.conv_blocks[0] | |
| cc = ConvCell0(device, blk.conv_b.weight, blk.conv_b.bias) | |
| assert cc.supports(xt) | |
| y = cc(xt) | |
| po = PoolCell0(device)(y) | |
| got = ttnn.to_torch(po).reshape(-1, 64).float() | |
| n_diff = int((got != ref).sum()) | |
| maxd = float((got - ref).abs().max()) | |
| print(f"cell0 pooled vs ttnn conv + pool: n_diff={n_diff} of {ref.numel()} max|diff|={maxd}") | |
| assert got.shape == ref.shape | |
| assert n_diff < 1e-4 * ref.numel() and maxd <= 0.125 | |
| cc.release() | |
| for t in (xt, y, po): | |
| ttnn.deallocate(t) | |
| tt_model.release() | |
| ttnn.deallocate(tt_in) | |
| def test_superpoint_cell1_conv_matches_torch(device): | |
| """Block-1 3x3 convs on the cell-tile layout (CellConv(G1), SP_CELL1=1): conv + bias + ReLU (and the | |
| fused horizontal pool max) against a torch fp32 conv of the same bf16 input / weights on all 120 | |
| cores (differences = bf16 output rounding of the HiFi2 / fp32-accumulated sums).""" | |
| from models.tt import conv_cell as C | |
| torch.manual_seed(0) | |
| g = C.G1 | |
| grid = ttnn.num_cores_to_corerangeset(120, device.compute_with_storage_grid_size(), row_wise=True) | |
| x = torch.relu(torch.randn(1, 1, 120 * g.cells, g.P * 64)).to(torch.bfloat16) | |
| xt = ttnn.from_torch(x, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device, | |
| memory_config=C.cell_memory_config(grid, g, g.P)) | |
| w = torch.randn(64, 64, 3, 3) * 0.05 | |
| b = torch.randn(64) * 0.1 | |
| img = x.float().reshape(240, 320, 64).permute(2, 0, 1)[None] | |
| yref = torch.relu(F.conv2d(img, w.to(torch.bfloat16).float(), b.to(torch.bfloat16).float(), padding=1))[0].permute(1, 2, 0) | |
| for hmax in (False, True): | |
| cc = C.CellConv(device, w, b, g, hmax=hmax) | |
| assert cc.supports(xt) | |
| y = cc(xt) | |
| got = ttnn.to_torch(y).float() | |
| ref = (torch.maximum(yref[:, 0::2], yref[:, 1::2]) if hmax else yref).reshape(got.shape[-2], -1) | |
| d = (got.reshape(ref.shape) - ref).abs() | |
| print(f"cell1 hmax={hmax}: max|d|={float(d.max()):.4g} mean|d|={float(d.mean()):.3g}") | |
| assert float(d.max()) < 0.05 and float(d.mean()) < 0.005 | |
| ttnn.deallocate(y) | |
| cc.release() | |
| ttnn.deallocate(xt) | |
| def test_superpoint_desc_head_matches_chain(device): | |
| """DescHeadRM (1x1 conv + L2 norm + untilize in one op, SP_DESC_HEAD=1) is bit-identical to the | |
| ttnn 1x1 conv (model head config) + DescNormRM on a [4800, 256] map sharded like the 3x3 head conv's | |
| output (64-row shards on the 12x10 grid, 75 used).""" | |
| from models.tt.desc_head import DescHeadRM | |
| from models.tt.desc_norm import DescNormRM | |
| torch_model = load_reference_model() | |
| dd = torch_model.descriptor_decoder | |
| w, b = dd.conv_descriptor_b.weight, dd.conv_descriptor_b.bias | |
| grid = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(11, 9))]) | |
| mc = ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, | |
| ttnn.ShardSpec(grid, [64, 256], ttnn.ShardOrientation.ROW_MAJOR)) | |
| g = torch.Generator().manual_seed(3) | |
| x = torch.relu(torch.randn(1, 1, 4800, 256, generator=g) * 0.5).to(torch.bfloat16) | |
| def up(): | |
| return ttnn.from_torch(x, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device, memory_config=mc) | |
| xt = up() | |
| dh = DescHeadRM(device, w, b, ttnn.L1_MEMORY_CONFIG) | |
| assert dh.supports(xt) | |
| got = ttnn.to_torch(dh(xt)).float().reshape(-1, 256) | |
| conv = TtConv2D(w, b, in_channels=256, out_channels=256, kernel_size=1, padding=0, device=device, activation=None, | |
| weights_dtype=ttnn.bfloat16, math_fidelity=ttnn.MathFidelity.HiFi2, fp32_dest_acc_en=True) | |
| z, _, _ = conv(up(), 60, 80, 1) | |
| ref = ttnn.to_torch(DescNormRM(device, ttnn.L1_MEMORY_CONFIG)(z)).float().reshape(-1, 256) | |
| n_diff = int((got != ref).sum()) | |
| print(f"desc head vs 1x1 conv + DescNormRM: n_diff={n_diff} of {ref.numel()}") | |
| assert n_diff == 0 | |
| def test_superpoint_merged_head_matches_chain(device, softmax): | |
| """SP_HEAD_MERGE: DescHeadRM on the merged [4800, 512] head conv output (descriptor channels 0..255, score | |
| 256..511) gives the descriptor map of ttnn 1x1 conv + DescNormRM on the first half AND the score logits of | |
| the ttnn score 1x1 conv (model head config) on the second half, bit for bit (incl. the softmax of both).""" | |
| from models.tt.desc_head import DescHeadRM | |
| from models.tt.desc_norm import DescNormRM | |
| torch_model = load_reference_model() | |
| dd, kd = torch_model.descriptor_decoder, torch_model.keypoint_decoder | |
| w, b = dd.conv_descriptor_b.weight, dd.conv_descriptor_b.bias | |
| ws, bs = kd.conv_score_b.weight, kd.conv_score_b.bias | |
| grid = ttnn.CoreRangeSet([ttnn.CoreRange(ttnn.CoreCoord(0, 0), ttnn.CoreCoord(11, 9))]) | |
| def mc(c): | |
| return ttnn.MemoryConfig(ttnn.TensorMemoryLayout.HEIGHT_SHARDED, ttnn.BufferType.L1, | |
| ttnn.ShardSpec(grid, [64, c], ttnn.ShardOrientation.ROW_MAJOR)) | |
| g = torch.Generator().manual_seed(4) | |
| x = torch.relu(torch.randn(1, 1, 4800, 512, generator=g) * 0.5).to(torch.bfloat16) | |
| def up(t): | |
| return ttnn.from_torch(t.contiguous(), dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=device, memory_config=mc(t.shape[-1])) | |
| dh = DescHeadRM(device, w, b, ttnn.L1_MEMORY_CONFIG, score_wb=(ws, bs), softmax=softmax) | |
| xt = up(x) | |
| assert dh.supports(xt) | |
| if softmax: | |
| assert dh.softmax_supported(xt) | |
| got_d = ttnn.to_torch(dh(xt)).float().reshape(-1, 256) | |
| s_got = dh.score_out | |
| got_s = ttnn.to_torch(s_got).float().reshape(-1, 65) | |
| # SP_HEAD_SM: the softmax computed inside the op on the idle cores; else ttnn.softmax on the logits | |
| got_sm = ttnn.to_torch(dh.smax_out if softmax else ttnn.softmax(s_got, dim=-1)).float().reshape(-1, 65) | |
| kw = dict(kernel_size=1, padding=0, device=device, activation=None, | |
| weights_dtype=ttnn.bfloat16, math_fidelity=ttnn.MathFidelity.HiFi2, fp32_dest_acc_en=True) | |
| conv_d = TtConv2D(w, b, in_channels=256, out_channels=256, **kw) | |
| conv_s = TtConv2D(ws, bs, in_channels=256, out_channels=65, **kw) | |
| z, _, _ = conv_d(up(x[..., :256]), 60, 80, 1) | |
| ref_d = ttnn.to_torch(DescNormRM(device, ttnn.L1_MEMORY_CONFIG)(z)).float().reshape(-1, 256) | |
| zs, _, _ = conv_s(up(x[..., 256:]), 60, 80, 1) | |
| ref_s = ttnn.to_torch(zs).float().reshape(-1, 65) | |
| ref_sm = ttnn.to_torch(ttnn.softmax(zs, dim=-1)).float().reshape(-1, 65) | |
| nd, ns, nsm = int((got_d != ref_d).sum()), int((got_s != ref_s).sum()), int((got_sm != ref_sm).sum()) | |
| print(f"merged head vs chain: desc n_diff={nd}, score logits n_diff={ns}, softmax n_diff={nsm}") | |
| assert nd == 0 and ns == 0 and nsm == 0 | |