romav2 2.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
romav2/__init__.py ADDED
@@ -0,0 +1,8 @@
1
+ import logging as _logging
2
+ from .logging import configure_logger as configure_logger
3
+ from .logging import logger as _logger
4
+
5
+ if not any(not isinstance(h, _logging.NullHandler) for h in _logger.handlers):
6
+ configure_logger()
7
+
8
+ from .romav2 import RoMaV2 as RoMaV2
@@ -0,0 +1,4 @@
1
+ from .scannet1500 import ScanNet1500 as ScanNet1500
2
+ from .mega1500 import Mega1500 as Mega1500
3
+ from .wxbs import WxBSBenchmark as WxBSBenchmark
4
+ from .satast import SatAst as SatAst
@@ -0,0 +1,115 @@
1
+ import numpy as np
2
+ import logging
3
+ from romav2.geometry import (
4
+ compute_pose_error,
5
+ pose_auc,
6
+ estimate_pose_cv2_ransac,
7
+ compute_relative_pose,
8
+ )
9
+ from PIL import Image
10
+ from tqdm import tqdm
11
+
12
+ logger = logging.getLogger(__name__)
13
+
14
+
15
+ class Mega1500:
16
+ def __init__(self, data_root="data/megadepth") -> None:
17
+ self.scene_names = [
18
+ "0015_0.1_0.3.npz",
19
+ "0015_0.3_0.5.npz",
20
+ "0022_0.1_0.3.npz",
21
+ "0022_0.3_0.5.npz",
22
+ "0022_0.5_0.7.npz",
23
+ ]
24
+ self.scenes = [
25
+ np.load(f"{data_root}/{scene}", allow_pickle=True)
26
+ for scene in self.scene_names
27
+ ]
28
+ self.data_root = data_root
29
+
30
+ def benchmark(self, model, model_name=None):
31
+ data_root = self.data_root
32
+ tot_e_t, tot_e_R, tot_e_pose = [], [], []
33
+ thresholds = [5, 10, 20]
34
+ for scene_ind in range(len(self.scenes)):
35
+ scene = self.scenes[scene_ind]
36
+ pairs = scene["pair_infos"]
37
+ intrinsics = scene["intrinsics"]
38
+ poses = scene["poses"]
39
+ im_paths = scene["image_paths"]
40
+ pair_inds = range(len(pairs))
41
+ for pairind in (pbar := tqdm(pair_inds, desc="Mega1500 eval")):
42
+ idx1, idx2 = pairs[pairind][0]
43
+ K1 = intrinsics[idx1].copy()
44
+ T1 = poses[idx1].copy()
45
+ R1, t1 = T1[:3, :3], T1[:3, 3]
46
+ K2 = intrinsics[idx2].copy()
47
+ T2 = poses[idx2].copy()
48
+ R2, t2 = T2[:3, :3], T2[:3, 3]
49
+ R, t = compute_relative_pose(R1, t1, R2, t2)
50
+ im_A_path = f"{data_root}/{im_paths[idx1]}"
51
+ im_B_path = f"{data_root}/{im_paths[idx2]}"
52
+ preds = model.match(im_A_path, im_B_path)
53
+ im_A = Image.open(im_A_path)
54
+ w1, h1 = im_A.size
55
+ im_B = Image.open(im_B_path)
56
+ w2, h2 = im_B.size
57
+ if True: # Note: we keep this true as it was used in DKM/RoMa papers. There is very little difference compared to setting to False.
58
+ scale1 = 1200 / max(w1, h1)
59
+ scale2 = 1200 / max(w2, h2)
60
+ w1, h1 = scale1 * w1, scale1 * h1
61
+ w2, h2 = scale2 * w2, scale2 * h2
62
+ K1, K2 = K1.copy(), K2.copy()
63
+ K1[:2] = K1[:2] * scale1
64
+ K2[:2] = K2[:2] * scale2
65
+ for _ in range(5):
66
+ matches, _, _, _ = model.sample(preds, 5_000)
67
+ kpts1, kpts2 = model.to_pixel_coordinates(matches, h1, w1, h2, w2)
68
+ kpts1, kpts2 = kpts1.cpu().numpy(), kpts2.cpu().numpy()
69
+ shuffling = np.random.permutation(np.arange(len(kpts1)))
70
+ kpts1 = kpts1[shuffling]
71
+ kpts2 = kpts2[shuffling]
72
+ try:
73
+ threshold = 0.5
74
+ norm_threshold = threshold / (
75
+ np.mean(np.abs(K1[:2, :2])) + np.mean(np.abs(K2[:2, :2]))
76
+ )
77
+ R_est, t_est, _ = estimate_pose_cv2_ransac(
78
+ kpts1,
79
+ kpts2,
80
+ K1,
81
+ K2,
82
+ norm_threshold,
83
+ conf=0.99999,
84
+ )
85
+ e_t, e_R = compute_pose_error(R_est, t_est[:, 0], R, t)
86
+ e_pose = max(e_t, e_R)
87
+ except Exception as e:
88
+ logger.debug(f"Pose estimation error: {e}")
89
+ e_t, e_R = 90, 90
90
+ e_pose = max(e_t, e_R)
91
+ tot_e_t.append(e_t)
92
+ tot_e_R.append(e_R)
93
+ tot_e_pose.append(e_pose)
94
+ pbar.set_postfix(
95
+ auc=f"{[f'{a.item():.3f}' for a in pose_auc(tot_e_pose, thresholds)]}"
96
+ )
97
+
98
+ tot_e_pose = np.array(tot_e_pose)
99
+ auc = pose_auc(tot_e_pose, thresholds)
100
+ acc_5 = (tot_e_pose < 5).mean()
101
+ acc_10 = (tot_e_pose < 10).mean()
102
+ acc_15 = (tot_e_pose < 15).mean()
103
+ acc_20 = (tot_e_pose < 20).mean()
104
+ map_5 = acc_5
105
+ map_10 = np.mean([acc_5, acc_10])
106
+ map_20 = np.mean([acc_5, acc_10, acc_15, acc_20])
107
+ logger.info("%s auc: %s", model_name, auc)
108
+ return {
109
+ "auc_5": auc[0],
110
+ "auc_10": auc[1],
111
+ "auc_20": auc[2],
112
+ "map_5": map_5,
113
+ "map_10": map_10,
114
+ "map_20": map_20,
115
+ }
@@ -0,0 +1,463 @@
1
+ import os
2
+ import logging
3
+ import torch
4
+ import tempfile
5
+ import json
6
+ import numpy as np
7
+ import cv2
8
+ from PIL import Image
9
+ from tqdm import tqdm
10
+ from romav2.geometry import pose_auc
11
+ import matplotlib.pyplot as plt
12
+
13
+ OFFSET = 0.5
14
+
15
+ logger = logging.getLogger(__name__)
16
+
17
+
18
+ class SatAst:
19
+ """
20
+ A benchmark class to evaluate image matchers using a folder of JSON files.
21
+
22
+ The final metric is the Area Under the Curve (AUC) of the cumulative
23
+ distribution of *all* individual errors, calculated at different
24
+ pixel thresholds.
25
+
26
+ Args:
27
+ json_folder_path (str): Path to the directory containing the .json files.
28
+ image_dataset_root (str): The root directory to prepend to the
29
+ image paths (query_path, pred_path)
30
+ found inside the JSON files.
31
+ """
32
+
33
+ def __init__(
34
+ self,
35
+ json_folder_path="data/satast/annotations-v1",
36
+ image_dataset_root="data/satast",
37
+ ) -> None:
38
+ self.json_folder_path = json_folder_path
39
+ self.image_dataset_root = image_dataset_root
40
+
41
+ try:
42
+ self.json_files = sorted(
43
+ [f for f in os.listdir(json_folder_path) if f.endswith(".json")]
44
+ )
45
+ if not self.json_files:
46
+ logger.warning("No .json files found in %s", json_folder_path)
47
+ except FileNotFoundError:
48
+ logger.error("Directory not found: %s", json_folder_path)
49
+ self.json_files = []
50
+
51
+ logger.info("Found %d JSON files.", len(self.json_files))
52
+
53
+ def _pixel_to_normalized(self, pts_pix, w, h, offset=OFFSET):
54
+ """
55
+ Converts pixel coordinates [0, n-1] to normalized [-1, 1].
56
+ """
57
+ pts_norm = np.zeros_like(pts_pix, dtype=np.float64)
58
+ pts_norm[..., 0] = (2.0 * (pts_pix[..., 0] + offset) / w) - 1.0
59
+ pts_norm[..., 1] = (2.0 * (pts_pix[..., 1] + offset) / h) - 1.0
60
+ return pts_norm
61
+
62
+ def _normalized_to_pixel(self, pts_norm, w, h, offset=OFFSET):
63
+ """
64
+ Converts normalized coordinates [-1, 1] to pixel [0, n-1].
65
+ """
66
+ pts_pix = np.zeros_like(pts_norm, dtype=np.float64)
67
+ pts_pix[..., 0] = (w * (pts_norm[..., 0] + 1.0) / 2.0) - offset
68
+ pts_pix[..., 1] = (h * (pts_norm[..., 1] + 1.0) / 2.0) - offset
69
+ return pts_pix
70
+
71
+ def benchmark(self, model, model_name=None, visualize=False, num_visualize=10):
72
+ """
73
+ Runs the full benchmark on the provided model.
74
+ This version benchmarks correspondences through an estimated homography.
75
+
76
+ Args:
77
+ model: The matcher model object. Must have a 'match(im_A_path, im_B_path)'
78
+ method and a 'sample(matches, certainty, num_samples)' method.
79
+ Assumed to return matches in normalized [-1, 1] coordinates.
80
+ model_name (str, optional): A name for the model.
81
+ visualize (bool, optional): If True, plots the first `num_visualize`
82
+ successful homography estimations.
83
+ num_visualize (int, optional): The maximum number of pairs to visualize.
84
+
85
+ Returns:
86
+ dict: A dictionary containing the AUC scores.
87
+ """
88
+ # This list will store *all* individual reprojection errors
89
+ all_reprojection_errors = []
90
+ viz_counter = 0
91
+ temp_file = tempfile.NamedTemporaryFile(suffix=".jpeg", delete=False)
92
+ im_B_path = temp_file.name
93
+
94
+ for json_name in tqdm(self.json_files, desc="Running Benchmark"):
95
+ json_path = os.path.join(self.json_folder_path, json_name)
96
+
97
+ # 1. Load JSON and Image Paths
98
+ with open(json_path, "r") as f:
99
+ json_data = json.load(f)
100
+
101
+ im_A_path_rel = json_data["query_path"].replace("\\", "/")
102
+ im_B_path_rel = json_data["pred_path"].replace("\\", "/")
103
+
104
+ im_A_path = os.path.join(self.image_dataset_root, im_A_path_rel)
105
+ im_B_path0 = os.path.join(self.image_dataset_root, im_B_path_rel)
106
+
107
+ im_A = Image.open(im_A_path)
108
+ w1, h1 = im_A.size
109
+ im_B0 = Image.open(im_B_path0)
110
+ # if "__rot0__" not in im_B_path0:
111
+ # continue
112
+ w2, h2 = im_B0.size
113
+
114
+ for rot_idx in range(1):
115
+ im_B0 = Image.open(im_B_path0)
116
+
117
+ if rot_idx == 0:
118
+ rotated_im_B = im_B0
119
+ elif rot_idx == 1:
120
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_90)
121
+ elif rot_idx == 2:
122
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_180)
123
+ elif rot_idx == 3:
124
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_270)
125
+ rotated_im_B.save(im_B_path, format="JPEG")
126
+
127
+ # 2. Run the matcher model
128
+ dense_preds = model.match(im_A_path, im_B_path)
129
+ if isinstance(dense_preds, dict):
130
+ sparse_preds = model.sample(dense_preds, 10_000)
131
+ else:
132
+ sparse_preds = model.sample(*dense_preds, 10_000)
133
+ good_matches = sparse_preds[0]
134
+ pos_a_norm = good_matches[:, :2]
135
+ pos_b_norm = good_matches[:, 2:]
136
+
137
+ # Rotate back
138
+ if rot_idx == 1:
139
+ pos_b_norm = torch.stack(
140
+ [-pos_b_norm[:, 1], pos_b_norm[:, 0]], dim=1
141
+ )
142
+ elif rot_idx == 2:
143
+ pos_b_norm = -pos_b_norm
144
+ elif rot_idx == 3:
145
+ pos_b_norm = torch.stack(
146
+ [pos_b_norm[:, 1], -pos_b_norm[:, 0]], dim=1
147
+ )
148
+
149
+ # 3. Estimate Homography in Normalized Coordinates
150
+ H_pred_norm = None
151
+ if len(pos_a_norm) >= 4:
152
+ try:
153
+ norm_thresh = 0.001
154
+
155
+ H_pred_norm, inliers = cv2.findHomography(
156
+ pos_a_norm.cpu().numpy(),
157
+ pos_b_norm.cpu().numpy(),
158
+ method=cv2.USAC_DEFAULT,
159
+ confidence=0.99999999,
160
+ maxIters=100_000,
161
+ ransacReprojThreshold=norm_thresh,
162
+ )
163
+ except cv2.error as e:
164
+ print(e)
165
+ H_pred_norm = None
166
+
167
+ # 4. Get Ground Truth Correspondences (Pixel)
168
+ last_iteration_corrs = json_data["correspondences"][-1]
169
+ pts_src_pix = np.array(last_iteration_corrs["pts_src"])
170
+ pts_dst_pix = np.array(last_iteration_corrs["pts_dst"])
171
+
172
+ if pts_src_pix.shape[0] == 0:
173
+ continue # No GT points, skip this pair
174
+
175
+ if H_pred_norm is None:
176
+ # If RANSAC fails, append 'inf' for all GT points
177
+ all_reprojection_errors.extend(
178
+ [float("inf")] * pts_src_pix.shape[0]
179
+ )
180
+ continue
181
+
182
+ if visualize and viz_counter < num_visualize:
183
+ logger.info("Visualizing pair: %s", json_name)
184
+
185
+ # 1. Create pixel-to-norm (M1) and norm-to-pixel (M2_inv) matrices
186
+
187
+ # M1: P1 -> N1
188
+ M1 = np.array(
189
+ [
190
+ [2.0 / w1, 0, (2.0 * OFFSET / w1) - 1.0],
191
+ [0, 2.0 / h1, (2.0 * OFFSET / h1) - 1.0],
192
+ [0, 0, 1.0],
193
+ ],
194
+ dtype=np.float64,
195
+ )
196
+
197
+ # M2_inv: N2 -> P2
198
+ M2_inv = np.array(
199
+ [
200
+ [w2 / 2.0, 0, w2 / 2.0 - OFFSET],
201
+ [0, h2 / 2.0, h2 / 2.0 - OFFSET],
202
+ [0, 0, 1.0],
203
+ ],
204
+ dtype=np.float64,
205
+ )
206
+
207
+ # 2. Convert H_pred_norm (N1 -> N2) to H_pix (P1 -> P2)
208
+ H_pix = M2_inv @ H_pred_norm @ M1
209
+
210
+ # 3. Invert H_pix to get (P2 -> P1) for warping Image 2 to 1
211
+ try:
212
+ H_pix_inv = np.linalg.inv(H_pix)
213
+ except np.linalg.LinAlgError:
214
+ logger.warning(
215
+ "Could not invert homography for %s. Skipping visualization.",
216
+ json_name,
217
+ )
218
+ continue # Skip this viz
219
+
220
+ # --- NEW: Extract and sample inliers for visualization ---
221
+ plot_pts_a = None
222
+ plot_pts_b = None
223
+ plot_colors = None
224
+
225
+ # Get inlier points from the RANSAC
226
+ inliers_mask = inliers.flatten().astype(bool)
227
+ inlier_pts_a_norm = pos_a_norm.cpu().numpy()[inliers_mask]
228
+ inlier_pts_b_norm = pos_b_norm.cpu().numpy()[inliers_mask]
229
+ num_inliers = inlier_pts_a_norm.shape[0]
230
+
231
+ if num_inliers > 0:
232
+ # Sample up to 10 inliers
233
+ num_to_sample = min(num_inliers, 10)
234
+ indices = np.random.choice(
235
+ num_inliers, num_to_sample, replace=False
236
+ )
237
+
238
+ sampled_pts_a_norm = inlier_pts_a_norm[indices]
239
+ sampled_pts_b_norm = inlier_pts_b_norm[indices]
240
+
241
+ # Convert to pixel coordinates for plotting
242
+ plot_pts_a = self._normalized_to_pixel(
243
+ sampled_pts_a_norm, w1, h1
244
+ )
245
+ plot_pts_b = self._normalized_to_pixel(
246
+ sampled_pts_b_norm, w2, h2
247
+ )
248
+
249
+ # Generate random colors for each pair
250
+ plot_colors = np.random.rand(num_to_sample, 3)
251
+
252
+ # 4. Load images with OpenCV for warping
253
+ img1_cv = cv2.imread(im_A_path) # Target
254
+ img2_cv = cv2.imread(im_B_path0) # Source
255
+
256
+ # 5. Warp image 2 to image 1's frame
257
+ warped_img2 = cv2.warpPerspective(img2_cv, H_pix_inv, (w1, h1))
258
+
259
+ # 6. Plot
260
+ img1_rgb = cv2.cvtColor(img1_cv, cv2.COLOR_BGR2RGB)
261
+ img2_rgb = cv2.cvtColor(img2_cv, cv2.COLOR_BGR2RGB)
262
+ warped_img2_rgb = cv2.cvtColor(warped_img2, cv2.COLOR_BGR2RGB)
263
+
264
+ fig, axes = plt.subplots(1, 3, figsize=(21, 7))
265
+ fig.suptitle(
266
+ f"Homography Visualization: {os.path.basename(im_B_path0)} -> {os.path.basename(im_A_path)}",
267
+ fontsize=16,
268
+ )
269
+
270
+ axes[0].imshow(img1_rgb)
271
+ axes[0].set_title("Image 1 (Target)")
272
+ axes[0].axis("off")
273
+ if plot_pts_a is not None:
274
+ axes[0].scatter(
275
+ plot_pts_a[:, 0],
276
+ plot_pts_a[:, 1],
277
+ s=40,
278
+ c=plot_colors,
279
+ marker="x",
280
+ )
281
+
282
+ axes[1].imshow(img2_rgb)
283
+ axes[1].set_title("Image 2 (Source)")
284
+ axes[1].axis("off")
285
+ if plot_pts_b is not None:
286
+ axes[1].scatter(
287
+ plot_pts_b[:, 0],
288
+ plot_pts_b[:, 1],
289
+ s=40,
290
+ c=plot_colors,
291
+ marker="x",
292
+ )
293
+
294
+ axes[2].imshow(warped_img2_rgb)
295
+ axes[2].set_title("Image 2 Warped to Image 1")
296
+ axes[2].axis("off")
297
+
298
+ plt.tight_layout()
299
+ plt.show()
300
+
301
+ viz_counter += 1
302
+
303
+ # 5. Normalize GT Source Points
304
+ pts_src_norm = self._pixel_to_normalized(pts_src_pix, w1, h1)
305
+
306
+ # 6. Warp Normalized GT Points
307
+ pts_src_norm_h = np.hstack(
308
+ (pts_src_norm, np.ones((pts_src_norm.shape[0], 1)))
309
+ )
310
+ warped_pts_norm_h = np.dot(pts_src_norm_h, H_pred_norm.T)
311
+
312
+ # De-homogenize
313
+ epsilon = 1e-8
314
+ warped_pts_norm = warped_pts_norm_h[:, :2] / (
315
+ warped_pts_norm_h[:, 2, np.newaxis] + epsilon
316
+ )
317
+
318
+ # 7. Convert Warped Points back to Pixel Coords
319
+ warped_pts_pix = self._normalized_to_pixel(warped_pts_norm, w2, h2)
320
+
321
+ # 8. Calculate and Store All Errors
322
+ errors = np.linalg.norm(warped_pts_pix - pts_dst_pix, axis=1)
323
+ all_reprojection_errors.extend(errors)
324
+
325
+ temp_file.close()
326
+
327
+ # 9. Compute AUC over *all* errors
328
+ thresholds = np.arange(1, 31)
329
+ auc = pose_auc(np.array(all_reprojection_errors), thresholds)
330
+ logger.info("=== AUC OF ERRORS ===")
331
+ logger.info("%s", auc)
332
+
333
+ results = {
334
+ "reprojection_auc_5px": auc[4],
335
+ "reprojection_auc_10px": auc[9],
336
+ "reprojection_auc_20px": auc[19],
337
+ "reprojection_auc_30px": auc[29],
338
+ }
339
+ logger.info("=== MAIN RESULTS ===")
340
+ logger.info("%s", results)
341
+
342
+ return results
343
+
344
+ def benchmark_warp(self, model, model_name=None, visualize=False, num_visualize=10):
345
+ """
346
+ Runs the full benchmark on the provided model.
347
+ This version benchmarks correspondences through an estimated dense warp.
348
+
349
+ Args:
350
+ model: The matcher model object. Must have a 'match(im_A_path, im_B_path)'
351
+ method and a 'sample(matches, certainty, num_samples)' method.
352
+ Assumed to return matches in normalized [-1, 1] coordinates.
353
+ model_name (str, optional): A name for the model.
354
+ visualize (bool, optional): If True, plots the first `num_visualize`
355
+ successful homography estimations.
356
+ num_visualize (int, optional): The maximum number of pairs to visualize.
357
+
358
+ Returns:
359
+ dict: A dictionary containing the AUC scores.
360
+ """
361
+ # This list will store *all* individual reprojection errors
362
+ all_reprojection_errors = []
363
+ temp_file = tempfile.NamedTemporaryFile(suffix=".jpeg", delete=False)
364
+ im_B_path = temp_file.name
365
+
366
+ for json_name in tqdm(self.json_files, desc="Running Benchmark"):
367
+ json_path = os.path.join(self.json_folder_path, json_name)
368
+
369
+ # 1. Load JSON and Image Paths
370
+ with open(json_path, "r") as f:
371
+ json_data = json.load(f)
372
+
373
+ im_A_path_rel = json_data["query_path"].replace("\\", "/")
374
+ im_B_path_rel = json_data["pred_path"].replace("\\", "/")
375
+
376
+ im_A_path = os.path.join(self.image_dataset_root, im_A_path_rel)
377
+ im_B_path0 = os.path.join(self.image_dataset_root, im_B_path_rel)
378
+
379
+ im_A = Image.open(im_A_path)
380
+ w1, h1 = im_A.size
381
+ im_B0 = Image.open(im_B_path0)
382
+ w2, h2 = im_B0.size
383
+
384
+ for rot_idx in range(4):
385
+ im_B0 = Image.open(im_B_path0)
386
+
387
+ if rot_idx == 0:
388
+ rotated_im_B = im_B0
389
+ elif rot_idx == 1:
390
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_90)
391
+ elif rot_idx == 2:
392
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_180)
393
+ elif rot_idx == 3:
394
+ rotated_im_B = im_B0.transpose(method=Image.Transpose.ROTATE_270)
395
+ rotated_im_B.save(im_B_path, format="JPEG")
396
+
397
+ # 2. Run the matcher model
398
+ preds = model.match(im_A_path, im_B_path)
399
+ warp = preds["warp_AB"]
400
+
401
+ # 3. Get Ground Truth Correspondences (Pixel)
402
+ last_iteration_corrs = json_data["correspondences"][-1]
403
+ pts_src_pix = np.array(last_iteration_corrs["pts_src"])
404
+ pts_dst_pix = np.array(last_iteration_corrs["pts_dst"])
405
+
406
+ if pts_src_pix.shape[0] == 0:
407
+ continue # No GT points, skip this pair
408
+
409
+ # 5. Normalize GT Source Points
410
+ pts_src_norm = torch.tensor(
411
+ self._pixel_to_normalized(pts_src_pix, w1, h1),
412
+ device=warp.device,
413
+ dtype=warp.dtype,
414
+ )
415
+
416
+ # 6. Warp Normalized GT Points
417
+ warp_src_to_pred = warp[0, :, : warp.shape[2] // 2, -2:]
418
+ warped_pts_norm = torch.nn.functional.grid_sample(
419
+ warp_src_to_pred.permute(2, 0, 1)[None],
420
+ pts_src_norm[None, None],
421
+ align_corners=False,
422
+ mode="bilinear",
423
+ )[0, :, 0].mT
424
+
425
+ # Rotate back
426
+ if rot_idx == 1:
427
+ warped_pts_norm = torch.stack(
428
+ [-warped_pts_norm[:, 1], warped_pts_norm[:, 0]], dim=1
429
+ )
430
+ elif rot_idx == 2:
431
+ warped_pts_norm = -warped_pts_norm
432
+ elif rot_idx == 3:
433
+ warped_pts_norm = torch.stack(
434
+ [warped_pts_norm[:, 1], -warped_pts_norm[:, 0]], dim=1
435
+ )
436
+
437
+ # 7. Convert Warped Points back to Pixel Coords
438
+ warped_pts_pix = self._normalized_to_pixel(
439
+ warped_pts_norm.cpu().numpy(), w2, h2
440
+ )
441
+
442
+ # 8. Calculate and Store All Errors
443
+ errors = np.linalg.norm(warped_pts_pix - pts_dst_pix, axis=1)
444
+ all_reprojection_errors.extend(errors)
445
+
446
+ temp_file.close()
447
+
448
+ # 9. Compute AUC over *all* errors
449
+ thresholds = np.arange(1, 31)
450
+ auc = pose_auc(np.array(all_reprojection_errors), thresholds)
451
+ logger.info("=== AUC OF ERRORS ===")
452
+ logger.info("%s", auc)
453
+
454
+ results = {
455
+ "reprojection_auc_5px": auc[4],
456
+ "reprojection_auc_10px": auc[9],
457
+ "reprojection_auc_20px": auc[19],
458
+ "reprojection_auc_30px": auc[29],
459
+ }
460
+ logger.info("=== MAIN RESULTS ===")
461
+ logger.info("%s", results)
462
+
463
+ return results