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 +8 -0
- romav2/benchmarks/__init__.py +4 -0
- romav2/benchmarks/mega1500.py +115 -0
- romav2/benchmarks/satast.py +463 -0
- romav2/benchmarks/scannet1500.py +125 -0
- romav2/benchmarks/wxbs.py +104 -0
- romav2/device.py +9 -0
- romav2/dpt.py +516 -0
- romav2/features.py +191 -0
- romav2/geometry.py +261 -0
- romav2/io.py +24 -0
- romav2/local_correlation.py +152 -0
- romav2/logging.py +97 -0
- romav2/matcher.py +207 -0
- romav2/normalizers.py +17 -0
- romav2/refiner.py +277 -0
- romav2/romav2.py +531 -0
- romav2/types.py +75 -0
- romav2/vis.py +36 -0
- romav2/vit/__init__.py +304 -0
- romav2/vit/attention.py +181 -0
- romav2/vit/block.py +293 -0
- romav2/vit/ffn_layers.py +83 -0
- romav2/vit/layer_scale.py +29 -0
- romav2/vit/patch_embed.py +94 -0
- romav2/vit/rms_norm.py +24 -0
- romav2/vit/rope.py +133 -0
- romav2/vit/rope_mixed.py +111 -0
- romav2/vit/utils.py +48 -0
- romav2-2.0.0.dist-info/METADATA +161 -0
- romav2-2.0.0.dist-info/RECORD +32 -0
- romav2-2.0.0.dist-info/WHEEL +4 -0
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,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
|