diff --git a/tests/test_cycle_filter.py b/tests/test_cycle_filter.py new file mode 100644 index 0000000..9607674 --- /dev/null +++ b/tests/test_cycle_filter.py @@ -0,0 +1,81 @@ +"""Tests for rotation cycle consistency filtering.""" + +import numpy as np +import pycolmap +from scipy.spatial.transform import Rotation + +from vidmap.mapper.native.extension import native +from vidmap.mapper.native.state import SolveState +from vidmap.mapper.stages.relative_pose.cycle_filter import filter_pairs_by_cycle_consistency + + +def _build_test_solve_state( + rotations: list[np.ndarray], corrupt_pairs: dict[tuple[int, int], np.ndarray] | None = None +) -> SolveState: + rec = pycolmap.Reconstruction() + rec.add_camera_with_trivial_rig(pycolmap.Camera.create_from_model_name(1, "PINHOLE", 500.0, 640, 480)) + sidecars = native.MappingSidecars() + graph = pycolmap.PoseGraph() + + num_images = len(rotations) + for i in range(1, num_images + 1): + rec.add_image_with_trivial_frame(pycolmap.Image(image_id=i, camera_id=1, name=f"frame_{i:04d}.jpg")) + sidecars.add_image(i, native.ImageData()) + + for i in range(1, num_images + 1): + for j in range(i + 1, num_images + 1): + if corrupt_pairs and (i, j) in corrupt_pairs: + R_rel = corrupt_pairs[(i, j)] + else: + R_rel = rotations[j - 1] @ rotations[i - 1].T + pair = native.PairData() + pair.has_relative_pose = True + sidecars.add_pair(pycolmap.image_pair_to_pair_id(i, j), pair) + edge = pycolmap.PoseGraphEdge() + edge.valid = True + edge.cam2_from_cam1 = pycolmap.Rigid3d(rotation=pycolmap.Rotation3d(R_rel)) + graph.add_edge(i, j, edge) + + return SolveState(rec, graph, sidecars) + + +def test_cycle_filter_consistent_graph(): + rots = [ + Rotation.from_euler("xyz", [0, 0, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [10, 0, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [10, 15, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [0, 15, 0], degrees=True).as_matrix(), + ] + state = _build_test_solve_state(rots) + + num_filtered = filter_pairs_by_cycle_consistency( + state, + min_triangles=2, + max_median_cycle_error_deg=10.0, + max_inconsistent_ratio=0.5, + triangle_error_threshold_deg=5.0, + ) + assert num_filtered == 0 + + +def test_cycle_filter_inconsistent_edge(): + rots = [ + Rotation.from_euler("xyz", [0, 0, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [10, 0, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [10, 15, 0], degrees=True).as_matrix(), + Rotation.from_euler("xyz", [0, 15, 0], degrees=True).as_matrix(), + ] + # Corrupt edge (1, 2) with 80 degree bogus rotation + corrupt_R = Rotation.from_euler("xyz", [80, 0, 0], degrees=True).as_matrix() + state = _build_test_solve_state(rots, corrupt_pairs={(1, 2): corrupt_R}) + + num_filtered = filter_pairs_by_cycle_consistency( + state, + min_triangles=2, + max_median_cycle_error_deg=20.0, + max_inconsistent_ratio=0.6, + triangle_error_threshold_deg=15.0, + ) + assert num_filtered == 1 + corrupt_pid = pycolmap.image_pair_to_pair_id(1, 2) + assert not state.pose_graph.is_valid(corrupt_pid) diff --git a/vidmap/mapper/options/view_graph.py b/vidmap/mapper/options/view_graph.py index fa0aa07..ef34e32 100644 --- a/vidmap/mapper/options/view_graph.py +++ b/vidmap/mapper/options/view_graph.py @@ -55,6 +55,11 @@ class MDRPOptions: ransac_max_iterations: Annotated[int, Field(gt=0)] = 50000 ransac_max_epipolar_error: Annotated[float, Field(gt=0, allow_inf_nan=False)] = 4.0 depth_stddev_multiplier: Annotated[float, Field(gt=0, allow_inf_nan=False)] = 1.0 + filter_cycle_inconsistent_pairs: bool = True + min_triangles_for_cycle_check: int = 3 + max_median_cycle_error_deg: float = 20.0 + max_inconsistent_cycle_ratio: float = 0.6 + triangle_error_threshold_deg: float = 15.0 @pydantic_dataclass(frozen=True, config=ConfigDict(extra="forbid", strict=True)) diff --git a/vidmap/mapper/stages/relative_pose/cycle_filter.py b/vidmap/mapper/stages/relative_pose/cycle_filter.py new file mode 100644 index 0000000..e18c3a6 --- /dev/null +++ b/vidmap/mapper/stages/relative_pose/cycle_filter.py @@ -0,0 +1,93 @@ +"""Rotation cycle consistency filtering for relative poses.""" + +from __future__ import annotations + +import logging + +import numpy as np +import pycolmap + +from vidmap.mapper.native.state import SolveState + +logger = logging.getLogger(__name__) + + +def filter_pairs_by_cycle_consistency( + state: SolveState, + *, + min_triangles: int = 3, + max_median_cycle_error_deg: float = 20.0, + max_inconsistent_ratio: float = 0.6, + triangle_error_threshold_deg: float = 15.0, +) -> int: + """Filter image pairs whose relative rotation is inconsistent with 3-cycles in the view graph. + + For each valid pair (i, j) with an estimated relative pose, we examine all mutual + neighbors k such that (i, k) and (k, j) are also valid pairs with relative poses + (forming a triangle i-k-j). For each triangle, the cycle error is the angular + distance between R_ji and R_jk * R_ki. + If the pair has at least min_triangles triangles and the median cycle error exceeds + max_median_cycle_error_deg (and the fraction of inconsistent triangles exceeds + max_inconsistent_ratio), the pair is deemed an outlier and marked invalid. + + Returns the number of pairs filtered. + """ + adj: dict[int, set[int]] = {} + rot_lookup: dict[tuple[int, int], pycolmap.Rotation3d] = {} + + def has_relative_pose(pair_id: int) -> bool: + return state.pose_graph.is_valid(pair_id) and state.pair_data(pair_id).has_relative_pose + + for pair_id in state.pair_order: + if not has_relative_pose(pair_id): + continue + id1, id2 = pycolmap.pair_id_to_image_pair(pair_id) + adj.setdefault(id1, set()).add(id2) + adj.setdefault(id2, set()).add(id1) + rotation = state.pose_graph.edges[pair_id].cam2_from_cam1.rotation + rot_lookup[(id1, id2)] = rotation + rot_lookup[(id2, id1)] = rotation.inverse() + + flagged_pairs = [] + for pair_id in state.pair_order: + if not has_relative_pose(pair_id): + continue + id1, id2 = pycolmap.pair_id_to_image_pair(pair_id) + common = adj.get(id1, set()) & adj.get(id2, set()) + if len(common) < min_triangles: + continue + + R_21 = rot_lookup[(id1, id2)] + errors = [] + for k in common: + R_k1 = rot_lookup[(id1, k)] + R_2k = rot_lookup[(k, id2)] + R_pred = R_2k * R_k1 + err_deg = np.rad2deg((R_pred.inverse() * R_21).angle()) + errors.append(err_deg) + + median_error = float(np.median(errors)) + inconsistent_ratio = float(np.mean(np.array(errors) > triangle_error_threshold_deg)) + + if median_error > max_median_cycle_error_deg and inconsistent_ratio > max_inconsistent_ratio: + flagged_pairs.append((pair_id, id1, id2, len(common), median_error, inconsistent_ratio)) + + for pair_id, id1, id2, n_tri, med_err, inc_ratio in flagged_pairs: + state.pose_graph.set_invalid_edge(pair_id) + logger.info( + "Filtered pair %d (%d <-> %d) by cycle consistency: %d triangles, median_error=%.1f deg, inconsistent=%.1f%%", + pair_id, + id1, + id2, + n_tri, + med_err, + inc_ratio * 100.0, + ) + + if flagged_pairs: + logger.warning( + "Cycle consistency filtering invalidated %d inconsistent pairs", + len(flagged_pairs), + ) + + return len(flagged_pairs) diff --git a/vidmap/mapper/stages/relative_pose/estimator.py b/vidmap/mapper/stages/relative_pose/estimator.py index 51c3085..619401a 100644 --- a/vidmap/mapper/stages/relative_pose/estimator.py +++ b/vidmap/mapper/stages/relative_pose/estimator.py @@ -20,6 +20,7 @@ from vidmap.mapper.replay.evidence.stages import capture_relative_pose_state, relative_pose_summary from vidmap.utils.logging import progress_bars_enabled +from .cycle_filter import filter_pairs_by_cycle_consistency from .mdrp import MDRPResult, ValidMDRPResult, estimate_mdrp_pose_for_pair from .native_options import build_inlier_threshold_options @@ -177,6 +178,23 @@ def estimate(self) -> RelativePoseResult: ) filtered_consecutive_pairs = current + if self.options.filter_cycle_inconsistent_pairs: + filter_pairs_by_cycle_consistency( + state, + min_triangles=self.options.min_triangles_for_cycle_check, + max_median_cycle_error_deg=self.options.max_median_cycle_error_deg, + max_inconsistent_ratio=self.options.max_inconsistent_cycle_ratio, + triangle_error_threshold_deg=self.options.triangle_error_threshold_deg, + ) + current = {pair_id for pair_id in consecutive_pair_ids if not state.pose_graph.is_valid(pair_id)} + newly_filtered = len(current - filtered_consecutive_pairs) + if newly_filtered: + logger.warning( + "%d consecutive pairs were filtered out by cycle consistency, continuing...", + newly_filtered, + ) + filtered_consecutive_pairs = current + summary = ( relative_pose_summary( state,