def to_homogeneous(points: torch.Tensor) -> torch.Tensor:
"""Append a 1 to each 2D point."""
if points.ndim != 2 or points.shape[1] != 2:
raise ValueError("Expected points with shape (N, 2).")
ones = torch.ones((points.shape[0], 1), dtype=points.dtype, device=points.device)
return torch.cat([points, ones], dim=1)
def from_homogeneous(points: torch.Tensor, eps: float = 1e-12) -> torch.Tensor:
"""Convert homogeneous points back to Euclidean coordinates."""
if points.ndim != 2 or points.shape[1] != 3:
raise ValueError("Expected homogeneous points with shape (N, 3).")
scale = points[:, 2:].clone()
scale = torch.where(scale.abs() < eps, torch.full_like(scale, eps), scale)
return points[:, :2] / scale
def apply_homography(points: torch.Tensor, H: torch.Tensor) -> torch.Tensor:
"""Apply a 3x3 homography to 2D points."""
if H.shape != (3, 3):
raise ValueError("Expected H with shape (3, 3).")
warped = to_homogeneous(points) @ H.T
return from_homogeneous(warped)
def normalize_points(points: torch.Tensor, eps: float = 1e-12) -> tuple[torch.Tensor, torch.Tensor]:
"""Normalize points so the centroid is at the origin and mean distance is sqrt(2)."""
centroid = points.mean(dim=0)
centered = points - centroid
mean_dist = torch.linalg.norm(centered, dim=1).mean().clamp(min=eps)
scale = math.sqrt(2.0) / mean_dist
T = torch.tensor(
[
[scale, 0.0, -scale * centroid[0]],
[0.0, scale, -scale * centroid[1]],
[0.0, 0.0, 1.0],
],
dtype=points.dtype,
device=points.device,
)
normalized = apply_homography(points, T)
return normalized, T
def estimate_homography_dlt(src_points: torch.Tensor, dst_points: torch.Tensor) -> torch.Tensor:
"""Estimate a homography with normalized DLT."""
if src_points.shape != dst_points.shape or src_points.shape[0] < 4:
raise ValueError("Need at least four matching point pairs.")
src_norm, T_src = normalize_points(src_points)
dst_norm, T_dst = normalize_points(dst_points)
rows = []
for (x, y), (u, v) in zip(src_norm, dst_norm):
rows.append(
torch.tensor(
[-x, -y, -1.0, 0.0, 0.0, 0.0, u * x, u * y, u],
dtype=src_points.dtype,
device=src_points.device,
)
)
rows.append(
torch.tensor(
[0.0, 0.0, 0.0, -x, -y, -1.0, v * x, v * y, v],
dtype=src_points.dtype,
device=src_points.device,
)
)
A = torch.stack(rows)
_, _, vh = torch.linalg.svd(A)
H_norm = vh[-1].reshape(3, 3)
H = torch.linalg.inv(T_dst) @ H_norm @ T_src
return H / torch.linalg.norm(H)
def compute_reprojection_error(H: torch.Tensor, src_points: torch.Tensor, dst_points: torch.Tensor) -> torch.Tensor:
"""Compute Euclidean reprojection error for each correspondence."""
projected = apply_homography(src_points, H)
return torch.linalg.norm(projected - dst_points, dim=1)
def ransac_homography(
src_points: torch.Tensor,
dst_points: torch.Tensor,
threshold: float,
num_iters: int,
) -> tuple[torch.Tensor, torch.Tensor, dict]:
"""Estimate a homography robustly with four-point RANSAC."""
if src_points.shape[0] < 4:
raise ValueError("RANSAC needs at least four correspondences.")
n = src_points.shape[0]
best_H = None
best_inliers = None
best_count = -1
best_mean_error = float("inf")
for _ in range(num_iters):
sample_idx = torch.randperm(n)[:4]
try:
candidate_H = estimate_homography_dlt(src_points[sample_idx], dst_points[sample_idx])
except RuntimeError:
continue
errors = compute_reprojection_error(candidate_H, src_points, dst_points)
inliers = errors < threshold
count = int(inliers.sum().item())
mean_error = float(errors[inliers].mean().item()) if count > 0 else float("inf")
if count > best_count or (count == best_count and mean_error < best_mean_error):
best_H = candidate_H
best_inliers = inliers
best_count = count
best_mean_error = mean_error
if best_H is None or best_inliers is None or int(best_inliers.sum().item()) < 4:
raise RuntimeError("RANSAC failed to find a valid homography.")
refined_H = estimate_homography_dlt(src_points[best_inliers], dst_points[best_inliers])
refined_errors = compute_reprojection_error(refined_H, src_points, dst_points)
diagnostics = {
"mean_inlier_error": float(refined_errors[best_inliers].mean().item()),
"mean_all_error": float(refined_errors.mean().item()),
"num_inliers": int(best_inliers.sum().item()),
"inlier_ratio": float(best_inliers.double().mean().item()),
}
return refined_H, best_inliers, diagnostics
def normalize_homography_scale(H: torch.Tensor) -> torch.Tensor:
"""Normalize a homography so it can be compared up to scale."""
if abs(float(H[-1, -1])) > 1e-12:
return H / H[-1, -1]
return H / torch.linalg.norm(H)
def make_grid(num_x: int = 6, num_y: int = 6, spacing: float = 1.0) -> torch.Tensor:
xs = torch.linspace(-2.5, 2.5, steps=num_x) * spacing
ys = torch.linspace(-2.0, 2.0, steps=num_y) * spacing
yy, xx = torch.meshgrid(ys, xs, indexing="ij")
return torch.stack([xx.reshape(-1), yy.reshape(-1)], dim=1)