#!/usr/bin/env python3 """ Pose Graph Optimization Utilities for ThreeDMatch Test Dataset This module provides utilities for performing pose graph optimization using Open3D to generate consistent per-frame poses from relative transformations. """ import os import numpy as np import logging from typing import Dict, List, Tuple, Optional, Any import open3d as o3d logger = logging.getLogger(__name__) def load_gt_info(gt_path: str) -> Dict[Tuple[int, int], np.ndarray]: """ Load information matrices from gt.info file. The gt.info file contains information matrices (covariance matrices) for each pair. Format: Each entry consists of 7 lines: - Line 1: source_id target_id - Lines 2-7: 6x6 information matrix (upper triangular, row-major) Args: gt_path: Path to directory containing gt.info Returns: Dictionary mapping (src_id, tgt_id) -> 6x6 information matrix """ info_file = os.path.join(gt_path, 'gt.info') if not os.path.exists(info_file): logger.warning(f"Information file not found: {info_file}, using identity information matrices") return {} info_matrices = {} try: with open(info_file, 'r') as f: content = f.readlines() i = 0 while i < len(content): # Parse header line: source_id target_id line = content[i].strip().split() if len(line) < 2: i += 1 continue src_id = int(line[0]) tgt_id = int(line[1]) # Read 6x6 information matrix (upper triangular, row-major) # We need to reconstruct the full symmetric matrix info_upper = [] for j in range(1, 7): # Lines 2-7 contain the matrix if i + j >= len(content): break row = [float(x) for x in content[i + j].strip().split()] info_upper.append(row) if len(info_upper) == 6: # Reconstruct symmetric matrix from upper triangular info_matrix = np.zeros((6, 6)) for row in range(6): for col in range(row, 6): col_idx = col - row if col_idx < len(info_upper[row]): info_matrix[row, col] = info_upper[row][col_idx] if row != col: info_matrix[col, row] = info_upper[row][col_idx] info_matrices[(src_id, tgt_id)] = info_matrix i += 7 # Move to next entry logger.info(f"Loaded {len(info_matrices)} information matrices from {info_file}") except Exception as e: logger.warning(f"Error loading information file {info_file}: {e}, using identity information matrices") return {} return info_matrices def load_gt_log(gt_path: str) -> Dict[Tuple[int, int], np.ndarray]: """ Load relative transformations from gt.log file. Format of gt.log: - Line 1: target_frame_id \t source_frame_id \t total_frame_count - Lines 2-5: 4x4 transformation matrix (4 rows, tab-separated values) - Then repeats for next pair... Note: The transformation T transforms points FROM source TO target. So T @ point_source = point_target Args: gt_path: Path to directory containing gt.log Returns: Dictionary mapping (src_id, tgt_id) -> 4x4 transformation matrix """ log_file = os.path.join(gt_path, 'gt.log') if not os.path.exists(log_file): raise FileNotFoundError(f"Ground truth log file not found: {log_file}") with open(log_file) as f: content = f.readlines() result = {} i = 0 while i < len(content): # Parse header line: target_frame_id \t source_frame_id \t total_frame_count header_line = content[i].strip() if not header_line: i += 1 continue header_parts = header_line.split("\t") if len(header_parts) < 2: i += 1 continue try: # Note: gt.log format is target_frame_id \t source_frame_id tgt_id = int(header_parts[0]) src_id = int(header_parts[1]) # total_frame_count = int(header_parts[2]) if len(header_parts) > 2 else None except (ValueError, IndexError) as e: logger.warning(f"Skipping malformed header line at index {i}: {header_line}, error: {e}") i += 1 continue # Read 4x4 transformation matrix (next 4 lines) if i + 4 >= len(content): logger.warning(f"Incomplete transformation matrix at index {i}, skipping") break trans = np.zeros([4, 4]) try: trans[0] = [float(x) for x in content[i + 1].strip().split("\t")[:4]] trans[1] = [float(x) for x in content[i + 2].strip().split("\t")[:4]] trans[2] = [float(x) for x in content[i + 3].strip().split("\t")[:4]] trans[3] = [float(x) for x in content[i + 4].strip().split("\t")[:4]] except (ValueError, IndexError) as e: logger.warning(f"Error parsing transformation matrix at index {i}: {e}, skipping") i += 5 continue # Store as (src_id, tgt_id) -> transformation # The transformation T transforms points FROM src TO tgt: point_tgt = T @ point_src result[(src_id, tgt_id)] = trans i += 5 # Move to next entry (header + 4 matrix rows) logger.info(f"Loaded {len(result)} transformations from {log_file}") return result def _initialize_poses_from_mst(gt_log: Dict[Tuple[int, int], np.ndarray], frame_ids: List[int], frame_id_to_index: Dict[int, int], anchor_frame_id: int = 0) -> List[np.ndarray]: """ Initialize poses using Minimum Spanning Tree (MST) from the anchor frame. Uses MST to find unique paths from anchor to all nodes, then propagates poses. Args: gt_log: Dictionary mapping (src_id, tgt_id) -> 4x4 transformation matrix frame_ids: List of all frame IDs frame_id_to_index: Mapping from frame_id to node index anchor_frame_id: Anchor frame ID Returns: List of initial 4x4 pose matrices for each node """ from collections import defaultdict import heapq num_frames = len(frame_ids) initial_poses = [None] * num_frames # Check if anchor frame exists if anchor_frame_id not in frame_id_to_index: logger.warning(f"Anchor frame {anchor_frame_id} not found in frame_ids. Available frames: {sorted(frame_ids)}") # Use the first available frame as anchor if frame_ids: anchor_frame_id = frame_ids[0] logger.info(f"Using first available frame {anchor_frame_id} as anchor instead") else: logger.error("No frames available, cannot initialize poses") return [np.eye(4) for _ in range(num_frames)] # Initialize anchor frame with identity anchor_idx = frame_id_to_index[anchor_frame_id] initial_poses[anchor_idx] = np.eye(4) logger.info(f"MST initialization: Using frame {anchor_frame_id} (index {anchor_idx}) as anchor") # Build weighted graph for MST # Weight = 1 for all edges (we just want connectivity, not actual weights) graph = defaultdict(list) edge_transforms = {} # Store transformations for each edge edges_added = 0 edges_skipped = 0 for (src_id, tgt_id), transform in gt_log.items(): src_idx = frame_id_to_index.get(src_id) tgt_idx = frame_id_to_index.get(tgt_id) if src_idx is not None and tgt_idx is not None: # Add edge in both directions for undirected graph graph[src_idx].append((tgt_idx, 1.0)) # weight = 1.0 graph[tgt_idx].append((src_idx, 1.0)) # Store transformations edge_transforms[(src_idx, tgt_idx)] = np.linalg.inv(transform) edge_transforms[(tgt_idx, src_idx)] = transform edges_added += 1 else: edges_skipped += 1 if edges_skipped <= 5: # Log first few skipped edges logger.debug(f"Skipping edge ({src_id}, {tgt_id}): src_idx={src_idx}, tgt_idx={tgt_idx}") logger.info(f"MST graph: Added {edges_added} edges, skipped {edges_skipped} edges") logger.info(f"MST graph: {len(graph)}/{num_frames} nodes have edges") # Check graph connectivity using BFS from collections import deque connectivity_visited = set() connectivity_queue = deque([anchor_idx]) connectivity_visited.add(anchor_idx) while connectivity_queue: node = connectivity_queue.popleft() for neighbor, _ in graph.get(node, []): if neighbor not in connectivity_visited: connectivity_visited.add(neighbor) connectivity_queue.append(neighbor) logger.info(f"Graph connectivity: {len(connectivity_visited)}/{num_frames} nodes reachable from anchor {anchor_frame_id}") if len(connectivity_visited) < num_frames: unreachable = set(range(num_frames)) - connectivity_visited logger.warning(f"Graph has disconnected components! {len(unreachable)} nodes unreachable from anchor") logger.warning(f"Unreachable node indices: {sorted(list(unreachable))[:10]}{'...' if len(unreachable) > 10 else ''}") # Check if anchor is connected if anchor_idx not in graph or len(graph[anchor_idx]) == 0: logger.warning(f"Anchor frame {anchor_frame_id} (index {anchor_idx}) has no edges in graph!") # Try to find a connected node for fid in frame_ids: idx = frame_id_to_index[fid] if idx in graph and len(graph[idx]) > 0: logger.info(f"Found connected frame {fid} (index {idx}), using as anchor instead") anchor_frame_id = fid anchor_idx = idx initial_poses[anchor_idx] = np.eye(4) break # Build MST using Prim's algorithm starting from anchor mst_edges = {} # Maps child -> (parent, transform_from_parent_to_child) visited = set() # Priority queue: (weight, current_node, parent_node, transform) pq = [(0.0, anchor_idx, anchor_idx, np.eye(4))] while pq and len(visited) < num_frames: weight, current, parent, transform = heapq.heappop(pq) if current in visited: continue visited.add(current) # Only add edge if current is not the anchor (anchor has no parent) if current != anchor_idx: mst_edges[current] = (parent, transform) # Add neighbors to queue for neighbor, edge_weight in graph.get(current, []): if neighbor not in visited: # Get transformation from current to neighbor transform_to_neighbor = edge_transforms.get((current, neighbor)) if transform_to_neighbor is not None: heapq.heappush(pq, (edge_weight, neighbor, current, transform_to_neighbor)) logger.info(f"MST construction: Visited {len(visited)}/{num_frames} nodes, MST has {len(mst_edges)} edges") # Propagate poses along MST from anchor to all nodes # Use BFS on MST to ensure we process nodes in order from collections import deque # Build MST adjacency list (directed from parent to child) mst_adjacency = defaultdict(list) for child, (parent, transform) in mst_edges.items(): mst_adjacency[parent].append((child, transform)) # Verify all nodes in MST are reachable nodes_in_mst = set([anchor_idx]) for child in mst_edges.keys(): nodes_in_mst.add(child) logger.info(f"MST nodes: {len(nodes_in_mst)} nodes in MST (including anchor)") # BFS to propagate poses queue = deque([anchor_idx]) poses_initialized = 1 # Anchor already initialized while queue: current_idx = queue.popleft() current_pose = initial_poses[current_idx] if current_pose is None: logger.warning(f"Processing node {current_idx} but its pose is None!") continue # Process children in MST children = mst_adjacency.get(current_idx, []) if children: logger.debug(f"Node {current_idx} has {len(children)} children in MST") for child_idx, transform in children: if initial_poses[child_idx] is None: # Compute child's pose using Open3D's convention: pose_tgt = pose_src @ T_edge # The transform stored in MST is T_current_child from gt.log # gt.log contains T_src_tgt that transforms POINTS from src to tgt: point_tgt = T @ point_src # # For Open3D pose graph: pose_tgt = pose_src @ T_edge # If T transforms points FROM src TO tgt, then T_edge = T (same transformation) # This is because Open3D's T_edge represents the relative pose transformation child_pose = current_pose @ transform initial_poses[child_idx] = child_pose poses_initialized += 1 queue.append(child_idx) else: logger.debug(f"Child {child_idx} already has pose initialized") logger.info(f"Pose propagation: Initialized {poses_initialized} poses via BFS") # Fill in any unvisited frames with identity (disconnected components) unvisited_count = 0 for i in range(num_frames): if initial_poses[i] is None: unvisited_count += 1 if unvisited_count <= 5: # Log first few unvisited frames logger.warning(f"Frame {frame_ids[i]} (index {i}) not reachable from anchor frame {anchor_frame_id} in MST, using identity pose") initial_poses[i] = np.eye(4) visited_count = sum(1 for p in initial_poses if p is not None) logger.info(f"MST initialization: {visited_count}/{num_frames} poses initialized from anchor {anchor_frame_id}") if unvisited_count > 0: logger.warning(f"MST initialization: {unvisited_count} frames were not reachable and initialized with identity") return initial_poses def build_pose_graph(gt_log: Dict[Tuple[int, int], np.ndarray], gt_info: Optional[Dict[Tuple[int, int], np.ndarray]] = None, anchor_frame_id: int = 0, use_smart_initialization: bool = True) -> Tuple[o3d.pipelines.registration.PoseGraph, List[int]]: """ Build an Open3D PoseGraph from relative transformations. Args: gt_log: Dictionary mapping (src_id, tgt_id) -> 4x4 transformation matrix gt_info: Optional dictionary mapping (src_id, tgt_id) -> 6x6 information matrix anchor_frame_id: Frame ID to use as anchor (fixed frame) use_smart_initialization: If True, use path-based initialization; if False, use identity Returns: Tuple of (Open3D PoseGraph object, list of frame IDs) """ pose_graph = o3d.pipelines.registration.PoseGraph() # Collect all unique frame IDs all_frame_ids = set() for (src_id, tgt_id) in gt_log.keys(): all_frame_ids.add(src_id) all_frame_ids.add(tgt_id) frame_ids = sorted(list(all_frame_ids)) num_frames = len(frame_ids) frame_id_to_index = {fid: idx for idx, fid in enumerate(frame_ids)} logger.info(f"Building pose graph with {num_frames} frames") logger.info(f"Frame IDs in gt.log: {frame_ids[:10]}{'...' if len(frame_ids) > 10 else ''} (showing first 10)") logger.info(f"Requested anchor frame: {anchor_frame_id}") # Validate anchor frame exists if anchor_frame_id not in frame_id_to_index: logger.warning(f"Anchor frame {anchor_frame_id} not found in gt.log frame IDs!") if frame_ids: anchor_frame_id = frame_ids[0] logger.info(f"Using first available frame {anchor_frame_id} as anchor instead") else: raise ValueError("No frame IDs found in gt.log, cannot build pose graph") # Initialize node poses if use_smart_initialization: logger.info("Using MST-based initialization for pose graph nodes") initial_poses = _initialize_poses_from_mst(gt_log, frame_ids, frame_id_to_index, anchor_frame_id) else: logger.info("Using identity initialization for pose graph nodes") initial_poses = [np.eye(4) for _ in range(num_frames)] # Create nodes with initial poses for i, frame_id in enumerate(frame_ids): pose_graph.nodes.append( o3d.pipelines.registration.PoseGraphNode(initial_poses[i]) ) # Add edges (relative transformations) for (src_id, tgt_id), transform in gt_log.items(): src_idx = frame_id_to_index.get(src_id) tgt_idx = frame_id_to_index.get(tgt_id) if src_idx is None or tgt_idx is None: logger.warning(f"Skipping edge ({src_id}, {tgt_id}): frame ID not found") continue # Get information matrix if available if gt_info is not None and (src_id, tgt_id) in gt_info: information = gt_info[(src_id, tgt_id)] else: # Use identity information matrix (uniform uncertainty) information = np.eye(6) # Add edge to pose graph # IMPORTANT: Open3D expects the edge transformation T such that: # pose_tgt = pose_src @ T_edge # So T_edge should transform poses FROM src TO tgt # # If gt.log contains T_src_tgt (transforms points from src to tgt), use it directly # If gt.log contains T_tgt_src (transforms points from tgt to src), we need to invert it pose_graph.edges.append( o3d.pipelines.registration.PoseGraphEdge( src_idx, tgt_idx, transform, # <-- EDGE TRANSFORMATION SET HERE (line 418) information, uncertain=False # Set to True if you want to mark uncertain edges ) ) # Debug: Log first few edge transformations to verify direction if len(pose_graph.edges) <= 3: logger.info(f"DEBUG Edge ({src_id}->{tgt_id}): transform translation={transform[:3, 3]}, " f"rotation trace={np.trace(transform[:3, :3]):.4f}") logger.info(f"DEBUG Edge transform matrix:\n{transform}") logger.info(f"Added {len(pose_graph.edges)} edges to pose graph") return pose_graph, frame_ids def compute_pose_graph_statistics(pose_graph: o3d.pipelines.registration.PoseGraph, gt_log: Dict[Tuple[int, int], np.ndarray], frame_ids: List[int], frame_id_to_index: Dict[int, int], anchor_frame_id: int = 0) -> Dict[str, Any]: """ Compute statistics about the pose graph optimization. Args: pose_graph: Optimized pose graph gt_log: Dictionary mapping (src_id, tgt_id) -> 4x4 transformation matrix frame_ids: List of frame IDs frame_id_to_index: Mapping from frame_id to node index anchor_frame_id: Anchor frame ID Returns: Dictionary containing various statistics """ stats = {} # Basic graph statistics stats['num_nodes'] = len(pose_graph.nodes) stats['num_edges'] = len(pose_graph.edges) stats['anchor_frame_id'] = anchor_frame_id # Compute edge errors (residuals) edge_errors = [] rotation_errors = [] translation_errors = [] for edge in pose_graph.edges: src_idx = edge.source_node_id tgt_idx = edge.target_node_id # Get poses from nodes src_pose = pose_graph.nodes[src_idx].pose tgt_pose = pose_graph.nodes[tgt_idx].pose # Compute relative pose from optimized poses relative_pose_optimized = np.linalg.inv(src_pose) @ tgt_pose # Get ground truth relative pose src_id = frame_ids[src_idx] tgt_id = frame_ids[tgt_idx] if (src_id, tgt_id) in gt_log: relative_pose_gt = gt_log[(src_id, tgt_id)] elif (tgt_id, src_id) in gt_log: relative_pose_gt = np.linalg.inv(gt_log[(tgt_id, src_id)]) else: continue # Compute error error_pose = np.linalg.inv(relative_pose_gt) @ relative_pose_optimized # Rotation error (angle in degrees) rotation_matrix = error_pose[:3, :3] trace = np.trace(rotation_matrix) trace = np.clip(trace, -1.0, 3.0) rotation_angle = np.arccos((trace - 1) / 2) rotation_error_deg = np.degrees(rotation_angle) rotation_errors.append(rotation_error_deg) # Translation error (in meters) translation_error = np.linalg.norm(error_pose[:3, 3]) translation_errors.append(translation_error) # Combined error (weighted) edge_errors.append({ 'src_id': src_id, 'tgt_id': tgt_id, 'rotation_error_deg': rotation_error_deg, 'translation_error_m': translation_error }) # Compute statistics if rotation_errors: stats['rotation_errors'] = { 'mean': np.mean(rotation_errors), 'std': np.std(rotation_errors), 'min': np.min(rotation_errors), 'max': np.max(rotation_errors), 'median': np.median(rotation_errors) } if translation_errors: stats['translation_errors'] = { 'mean': np.mean(translation_errors), 'std': np.std(translation_errors), 'min': np.min(translation_errors), 'max': np.max(translation_errors), 'median': np.median(translation_errors) } # Graph connectivity statistics adjacency = {} for edge in pose_graph.edges: src_idx = edge.source_node_id tgt_idx = edge.target_node_id if src_idx not in adjacency: adjacency[src_idx] = [] adjacency[src_idx].append(tgt_idx) if tgt_idx not in adjacency: adjacency[tgt_idx] = [] adjacency[tgt_idx].append(src_idx) node_degrees = [len(adjacency.get(i, [])) for i in range(len(pose_graph.nodes))] stats['connectivity'] = { 'mean_degree': np.mean(node_degrees) if node_degrees else 0, 'min_degree': np.min(node_degrees) if node_degrees else 0, 'max_degree': np.max(node_degrees) if node_degrees else 0, 'isolated_nodes': sum(1 for d in node_degrees if d == 0) } # Check for loop closures (nodes with degree > 2) loop_closure_nodes = sum(1 for d in node_degrees if d > 2) stats['loop_closures'] = { 'num_loop_closure_nodes': loop_closure_nodes, 'has_loop_closures': loop_closure_nodes > 0 } # Pose spread statistics (how far nodes are from anchor) anchor_idx = frame_id_to_index.get(anchor_frame_id, 0) anchor_pose = pose_graph.nodes[anchor_idx].pose distances_from_anchor = [] for i, node in enumerate(pose_graph.nodes): if i == anchor_idx: continue node_pose = node.pose # Compute relative pose from anchor relative_pose = np.linalg.inv(anchor_pose) @ node_pose # Distance is translation magnitude distance = np.linalg.norm(relative_pose[:3, 3]) distances_from_anchor.append(distance) if distances_from_anchor: stats['pose_spread'] = { 'mean_distance_from_anchor': np.mean(distances_from_anchor), 'max_distance_from_anchor': np.max(distances_from_anchor), 'min_distance_from_anchor': np.min(distances_from_anchor), 'std_distance_from_anchor': np.std(distances_from_anchor) } stats['edge_errors'] = edge_errors return stats def visualize_initial_poses(gt_path: str, fragments_path: str, frame_ids: List[int], initial_poses: List[np.ndarray], anchor_frame_id: int = 0, num_fragments_to_visualize: int = 5, max_points_per_fragment: int = 10000) -> None: """ Visualize point clouds transformed by initial poses to debug pose graph initialization. Args: gt_path: Path to directory containing gt.log fragments_path: Path to directory containing fragment PLY files frame_ids: List of frame IDs (fragment IDs) initial_poses: List of initial 4x4 pose matrices anchor_frame_id: Anchor frame ID num_fragments_to_visualize: Number of fragments to visualize (randomly selected) max_points_per_fragment: Maximum points to load per fragment for visualization """ import random import glob logger.info("=" * 60) logger.info("VISUALIZING INITIAL POSES FOR DEBUG") logger.info("=" * 60) # Find fragment files fragment_files = sorted(glob.glob(os.path.join(fragments_path, "*.ply"))) if not fragment_files: logger.error(f"No fragment files found in {fragments_path}") return # Create mapping from frame_id to fragment file frame_id_to_file = {} for fragment_file in fragment_files: fragment_name = os.path.splitext(os.path.basename(fragment_file))[0] # Extract fragment ID (e.g., "cloud_bin_0" -> 0) try: if 'bin' in fragment_name: parts = fragment_name.split('_') if len(parts) >= 3 and parts[-2] == 'bin': fragment_id = int(parts[-1]) frame_id_to_file[fragment_id] = fragment_file except (ValueError, IndexError): continue if not frame_id_to_file: logger.error("Could not map frame IDs to fragment files") return # Select fragments to visualize (include anchor + random selection) available_frame_ids = [fid for fid in frame_ids if fid in frame_id_to_file] if anchor_frame_id in available_frame_ids: selected_frame_ids = [anchor_frame_id] remaining = [fid for fid in available_frame_ids if fid != anchor_frame_id] if len(remaining) >= num_fragments_to_visualize - 1: selected_frame_ids.extend(random.sample(remaining, num_fragments_to_visualize - 1)) else: selected_frame_ids.extend(remaining) else: if len(available_frame_ids) >= num_fragments_to_visualize: selected_frame_ids = random.sample(available_frame_ids, num_fragments_to_visualize) else: selected_frame_ids = available_frame_ids logger.info(f"Visualizing {len(selected_frame_ids)} fragments: {selected_frame_ids}") # Load and transform point clouds point_clouds = [] colors = [] # Color map for different fragments color_map = np.array([ [1.0, 0.0, 0.0], # Red for anchor [0.0, 1.0, 0.0], # Green [0.0, 0.0, 1.0], # Blue [1.0, 1.0, 0.0], # Yellow [1.0, 0.0, 1.0], # Magenta [0.0, 1.0, 1.0], # Cyan [0.5, 0.5, 0.5], # Gray [1.0, 0.5, 0.0], # Orange ]) frame_id_to_index = {fid: idx for idx, fid in enumerate(frame_ids)} for i, frame_id in enumerate(selected_frame_ids): if frame_id not in frame_id_to_index: logger.warning(f"Frame ID {frame_id} not found in frame_ids list") continue idx = frame_id_to_index[frame_id] pose = initial_poses[idx] fragment_file = frame_id_to_file[frame_id] # Load point cloud try: pcd = o3d.io.read_point_cloud(fragment_file) points = np.asarray(pcd.points) if len(points) == 0: logger.warning(f"Empty point cloud for fragment {frame_id}") continue # Downsample if too many points if len(points) > max_points_per_fragment: indices = np.random.choice(len(points), max_points_per_fragment, replace=False) points = points[indices] # Transform points using initial pose # If pose transforms from world to fragment, we need inv(pose) to transform points # Actually, if pose is the pose of the fragment in world coordinates, # then to transform points from fragment to world: points_world = pose @ points_fragment points_homogeneous = np.hstack([points, np.ones((len(points), 1))]) points_transformed = (pose @ points_homogeneous.T).T[:, :3] # Create Open3D point cloud pcd_transformed = o3d.geometry.PointCloud() pcd_transformed.points = o3d.utility.Vector3dVector(points_transformed) # Assign color color = color_map[i % len(color_map)] pcd_transformed.paint_uniform_color(color) point_clouds.append(pcd_transformed) # Log pose info translation = pose[:3, 3] logger.info(f"Fragment {frame_id} (index {idx}): translation={translation}, " f"points={len(points_transformed)}, " f"{'ANCHOR' if frame_id == anchor_frame_id else ''}") except Exception as e: logger.error(f"Error loading fragment {frame_id} from {fragment_file}: {e}") continue if not point_clouds: logger.error("No point clouds loaded for visualization") return # Visualize logger.info(f"\nVisualizing {len(point_clouds)} point clouds...") logger.info("Red = Anchor frame, other colors = other fragments") logger.info("If poses are correct, fragments should align spatially") o3d.visualization.draw_geometries( point_clouds, window_name=f"Initial Poses Debug (Anchor: {anchor_frame_id})", width=1920, height=1080 ) logger.info("Visualization closed") def debug_initial_poses(gt_path: str, fragments_path: str, anchor_frame_id: int = 0, num_fragments_to_visualize: int = 5, use_smart_initialization: bool = True, visualize_optimized: bool = True, only_optimized: bool = False) -> None: """ Debug function to visualize initial poses from pose graph initialization and optionally optimized poses after PGO. Args: gt_path: Path to directory containing gt.log fragments_path: Path to directory containing fragment PLY files anchor_frame_id: Anchor frame ID num_fragments_to_visualize: Number of fragments to visualize use_smart_initialization: Whether to use MST initialization or identity visualize_optimized: Whether to also visualize optimized poses after PGO only_optimized: If True, skip initial poses visualization and only show optimized poses """ # Load ground truth gt_log = load_gt_log(gt_path) gt_info = load_gt_info(gt_path) # Collect frame IDs all_frame_ids = set() for (src_id, tgt_id) in gt_log.keys(): all_frame_ids.add(src_id) all_frame_ids.add(tgt_id) frame_ids = sorted(list(all_frame_ids)) frame_id_to_index = {fid: idx for idx, fid in enumerate(frame_ids)} logger.info(f"Found {len(frame_ids)} frames in gt.log") # Get initial poses (needed for PGO even if not visualizing) if use_smart_initialization: logger.info("Computing initial poses using MST...") initial_poses = _initialize_poses_from_mst(gt_log, frame_ids, frame_id_to_index, anchor_frame_id) else: logger.info("Using identity poses...") initial_poses = [np.eye(4) for _ in range(len(frame_ids))] # Visualize initial poses (unless only_optimized is True) if not only_optimized: logger.info("\n" + "=" * 60) logger.info("VISUALIZING INITIAL POSES") logger.info("=" * 60) visualize_initial_poses( gt_path=gt_path, fragments_path=fragments_path, frame_ids=frame_ids, initial_poses=initial_poses, anchor_frame_id=anchor_frame_id, num_fragments_to_visualize=num_fragments_to_visualize ) # If requested, also visualize optimized poses if visualize_optimized or only_optimized: logger.info("\n" + "=" * 60) logger.info("RUNNING POSE GRAPH OPTIMIZATION") logger.info("=" * 60) # Build and optimize pose graph pose_graph, _ = build_pose_graph( gt_log=gt_log, gt_info=gt_info, anchor_frame_id=anchor_frame_id, use_smart_initialization=use_smart_initialization ) # Optimize optimize_pose_graph(pose_graph, method="LM", max_iterations=100) # Extract optimized poses optimized_poses = extract_per_frame_poses(pose_graph, frame_ids, anchor_frame_id) # Convert to list format for visualization optimized_poses_list = [optimized_poses.get(fid, np.eye(4)) for fid in frame_ids] logger.info("\n" + "=" * 60) logger.info("VISUALIZING OPTIMIZED POSES AFTER PGO") logger.info("=" * 60) visualize_initial_poses( gt_path=gt_path, fragments_path=fragments_path, frame_ids=frame_ids, initial_poses=optimized_poses_list, anchor_frame_id=anchor_frame_id, num_fragments_to_visualize=num_fragments_to_visualize ) # Compute and print statistics stats = compute_pose_graph_statistics( pose_graph=pose_graph, gt_log=gt_log, frame_ids=frame_ids, frame_id_to_index=frame_id_to_index, anchor_frame_id=anchor_frame_id ) print_pose_graph_statistics(stats) def print_pose_graph_statistics(stats: Dict[str, Any]) -> None: """ Print pose graph optimization statistics in a readable format. Args: stats: Statistics dictionary from compute_pose_graph_statistics """ logger.info("=" * 60) logger.info("POSE GRAPH OPTIMIZATION STATISTICS") logger.info("=" * 60) # Basic graph info logger.info(f"Graph Structure:") logger.info(f" Number of nodes (frames): {stats['num_nodes']}") logger.info(f" Number of edges (constraints): {stats['num_edges']}") logger.info(f" Anchor frame ID: {stats['anchor_frame_id']}") # Connectivity if 'connectivity' in stats: conn = stats['connectivity'] logger.info(f"\nConnectivity:") logger.info(f" Mean degree: {conn['mean_degree']:.2f}") logger.info(f" Min degree: {conn['min_degree']}") logger.info(f" Max degree: {conn['max_degree']}") logger.info(f" Isolated nodes: {conn['isolated_nodes']}") # Loop closures if 'loop_closures' in stats: lc = stats['loop_closures'] logger.info(f"\nLoop Closures:") logger.info(f" Nodes with loop closures (degree > 2): {lc['num_loop_closure_nodes']}") logger.info(f" Has loop closures: {lc['has_loop_closures']}") # Rotation errors if 'rotation_errors' in stats: rot_err = stats['rotation_errors'] logger.info(f"\nRotation Errors (degrees):") logger.info(f" Mean: {rot_err['mean']:.4f}") logger.info(f" Std: {rot_err['std']:.4f}") logger.info(f" Min: {rot_err['min']:.4f}") logger.info(f" Max: {rot_err['max']:.4f}") logger.info(f" Median: {rot_err['median']:.4f}") # Translation errors if 'translation_errors' in stats: trans_err = stats['translation_errors'] logger.info(f"\nTranslation Errors (meters):") logger.info(f" Mean: {trans_err['mean']:.6f}") logger.info(f" Std: {trans_err['std']:.6f}") logger.info(f" Min: {trans_err['min']:.6f}") logger.info(f" Max: {trans_err['max']:.6f}") logger.info(f" Median: {trans_err['median']:.6f}") # Pose spread if 'pose_spread' in stats: spread = stats['pose_spread'] logger.info(f"\nPose Spread (distance from anchor):") logger.info(f" Mean distance: {spread['mean_distance_from_anchor']:.4f} m") logger.info(f" Max distance: {spread['max_distance_from_anchor']:.4f} m") logger.info(f" Min distance: {spread['min_distance_from_anchor']:.4f} m") logger.info(f" Std distance: {spread['std_distance_from_anchor']:.4f} m") logger.info("=" * 60) def optimize_pose_graph(pose_graph: o3d.pipelines.registration.PoseGraph, max_iterations: int = 100, method: str = "LM") -> o3d.pipelines.registration.PoseGraph: """ Optimize pose graph using Open3D's global optimization. Args: pose_graph: Input pose graph max_iterations: Maximum number of iterations method: Optimization method ("LM" for Levenberg-Marquardt or "GN" for Gauss-Newton) Returns: Optimized pose graph """ # Set optimization options option = o3d.pipelines.registration.GlobalOptimizationOption( max_correspondence_distance=0.05, edge_prune_threshold=0.25, preference_loop_closure=0.1, reference_node=0, # First node as reference ) # Set convergence criteria with max_iterations # Note: GlobalOptimizationConvergenceCriteria doesn't accept keyword arguments # Use default constructor and set attributes if available convergence_criteria = o3d.pipelines.registration.GlobalOptimizationConvergenceCriteria() # Try to set attributes if they exist (Open3D version dependent) try: if hasattr(convergence_criteria, 'max_iteration'): convergence_criteria.max_iteration = max_iterations if hasattr(convergence_criteria, 'min_relative_increment'): convergence_criteria.min_relative_increment = 1e-6 if hasattr(convergence_criteria, 'min_relative_residual_increment'): convergence_criteria.min_relative_residual_increment = 1e-6 if hasattr(convergence_criteria, 'min_right_term'): convergence_criteria.min_right_term = 1e-6 except AttributeError: # If attributes can't be set, use default criteria logger.debug("Using default convergence criteria (attributes not settable)") # Perform global optimization if method == "LM": o3d.pipelines.registration.global_optimization( pose_graph, o3d.pipelines.registration.GlobalOptimizationLevenbergMarquardt(), convergence_criteria, option ) else: # GN o3d.pipelines.registration.global_optimization( pose_graph, o3d.pipelines.registration.GlobalOptimizationGaussNewton(), convergence_criteria, option ) logger.info(f"Pose graph optimization completed (max_iterations: {max_iterations})") return pose_graph def extract_per_frame_poses(pose_graph: o3d.pipelines.registration.PoseGraph, frame_ids: List[int], anchor_frame_id: int = 0) -> Dict[int, np.ndarray]: """ Extract per-frame poses from optimized pose graph. Args: pose_graph: Optimized pose graph frame_ids: List of frame IDs corresponding to nodes anchor_frame_id: Frame ID used as anchor Returns: Dictionary mapping frame_id -> 4x4 pose matrix (in anchor frame's coordinate system) """ per_frame_poses = {} anchor_idx = frame_ids.index(anchor_frame_id) if anchor_frame_id in frame_ids else 0 for idx, frame_id in enumerate(frame_ids): if idx < len(pose_graph.nodes): # Get pose from pose graph node pose_matrix = pose_graph.nodes[idx].pose # If this is not the anchor frame, the pose is already relative to anchor # (since we set anchor as reference node) per_frame_poses[frame_id] = pose_matrix.copy() else: logger.warning(f"Frame {frame_id} (index {idx}) not found in pose graph") # Use identity as fallback per_frame_poses[frame_id] = np.eye(4) # Ensure anchor frame has identity pose per_frame_poses[anchor_frame_id] = np.eye(4) logger.info(f"Extracted {len(per_frame_poses)} per-frame poses") return per_frame_poses def optimize_threedmatch_poses(gt_path: str, anchor_frame_id: int = 0, max_iterations: int = 100, method: str = "LM", use_smart_initialization: bool = True, print_statistics: bool = True) -> Dict[int, np.ndarray]: """ Complete pipeline: Load gt.log and gt.info, build pose graph, optimize, and extract poses. Args: gt_path: Path to directory containing gt.log and gt.info anchor_frame_id: Frame ID to use as anchor (default: 0) max_iterations: Maximum iterations for optimization method: Optimization method ("LM" or "GN") use_smart_initialization: If True, use path-based initialization; if False, use identity print_statistics: If True, print detailed optimization statistics Returns: Dictionary mapping frame_id -> 4x4 pose matrix """ # Load ground truth data logger.info(f"Loading ground truth data from {gt_path}") gt_log = load_gt_log(gt_path) gt_info = load_gt_info(gt_path) if not gt_log: raise ValueError(f"No ground truth transformations found in {gt_path}") # Build pose graph logger.info("Building pose graph...") pose_graph, frame_ids = build_pose_graph(gt_log, gt_info, anchor_frame_id, use_smart_initialization) frame_id_to_index = {fid: idx for idx, fid in enumerate(frame_ids)} # Compute and print initial statistics if print_statistics: logger.info("\nComputing initial pose graph statistics...") initial_stats = compute_pose_graph_statistics( pose_graph, gt_log, frame_ids, frame_id_to_index, anchor_frame_id ) logger.info("Initial State (before optimization):") print_pose_graph_statistics(initial_stats) # Optimize pose graph logger.info(f"\nOptimizing pose graph (method: {method}, max_iterations: {max_iterations})...") optimized_pose_graph = optimize_pose_graph(pose_graph, max_iterations, method) # Compute and print final statistics if print_statistics: logger.info("\nComputing optimized pose graph statistics...") final_stats = compute_pose_graph_statistics( optimized_pose_graph, gt_log, frame_ids, frame_id_to_index, anchor_frame_id ) logger.info("Final State (after optimization):") print_pose_graph_statistics(final_stats) # Print improvement summary if 'rotation_errors' in initial_stats and 'rotation_errors' in final_stats: rot_improvement = initial_stats['rotation_errors']['mean'] - final_stats['rotation_errors']['mean'] logger.info(f"\nOptimization Improvement:") logger.info(f" Rotation error reduction: {rot_improvement:.4f} degrees") if 'translation_errors' in initial_stats and 'translation_errors' in final_stats: trans_improvement = initial_stats['translation_errors']['mean'] - final_stats['translation_errors']['mean'] logger.info(f" Translation error reduction: {trans_improvement:.6f} meters") # Extract per-frame poses logger.info("\nExtracting per-frame poses...") per_frame_poses = extract_per_frame_poses(optimized_pose_graph, frame_ids, anchor_frame_id) return per_frame_poses