RAP / dataset_process /utils /pose_graph_utils.py
YuePanEdward's picture
Squash history: release superseded example-data blobs
be88765
Raw History Blame Contribute Delete
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