Spaces:
Running on Zero
Running on Zero
Download dataset_process/utils/pose_graph_utils.py from YuePanEdward/RAP: direct link, hf CLI and curl.
- Browser
- Download file 44.8 kB
-
https://huggingface.co/spaces/YuePanEdward/RAP/resolve/main/dataset_process/utils/pose_graph_utils.py
- Command line
-
hf download hf://spaces/YuePanEdward/RAP/dataset_process/utils/pose_graph_utils.py
-
curl -L -o pose_graph_utils.py https://huggingface.co/spaces/YuePanEdward/RAP/resolve/main/dataset_process/utils/pose_graph_utils.py
44.8 kB
| #!/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 | |