From 2d6a2d51fd18ed211097c170ad3eefaa9cf3e234 Mon Sep 17 00:00:00 2001 From: mhubii Date: Fri, 24 Jul 2026 19:02:56 +0100 Subject: [PATCH 01/12] data structure refactor --- cli/rr_cam_swarm.py | 16 +++-- cli/rr_hydra.py | 16 ++--- cli/rr_mono_dr.py | 22 ++++--- cli/rr_stereo_dr.py | 64 ++++++++++++------- roboreg/core/structs.py | 19 +++++- roboreg/io/meshes.py | 9 +-- roboreg/io/parsers.py | 92 +++++++-------------------- roboreg/io/robot_data.py | 22 +------ roboreg/reg/__init__.py | 0 roboreg/reg/_validation.py | 43 +++++++++++++ roboreg/reg/img/__init__.py | 0 roboreg/reg/img/request.py | 96 ++++++++++++++++++++++++++++ roboreg/reg/pcl/__init__.py | 0 roboreg/reg/pcl/request.py | 49 +++++++++++++++ test/io/test_parsers.py | 122 ++++++++++++++++++++---------------- test/test_hydra_icp.py | 30 +++++---- 16 files changed, 388 insertions(+), 212 deletions(-) create mode 100644 roboreg/reg/__init__.py create mode 100644 roboreg/reg/_validation.py create mode 100644 roboreg/reg/img/__init__.py create mode 100644 roboreg/reg/img/request.py create mode 100644 roboreg/reg/pcl/__init__.py create mode 100644 roboreg/reg/pcl/request.py diff --git a/cli/rr_cam_swarm.py b/cli/rr_cam_swarm.py index 1a52e93..e74fbc7 100644 --- a/cli/rr_cam_swarm.py +++ b/cli/rr_cam_swarm.py @@ -19,7 +19,7 @@ load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, parse_camera_info, - parse_mono_data, + parse_monocular_observations, ) from roboreg.losses import soft_dice_loss from roboreg.optim import LinearParticleSwarm, ParticleSwarmOptimizer @@ -274,7 +274,7 @@ def main() -> None: image_files = np.array(image_files)[random_indices].tolist() mask_files = np.array(mask_files)[random_indices].tolist() joint_states_files = np.array(joint_states_files)[random_indices].tolist() - images, joint_states, masks = parse_mono_data( + observations = parse_monocular_observations( image_files=image_files, mask_files=mask_files, joint_states_files=joint_states_files, @@ -282,10 +282,10 @@ def main() -> None: # pre-process data joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) n_joint_states = joint_states.shape[0] - masks = [mask_exponential_decay(mask) for mask in masks] + masks = [mask_exponential_decay(mask) for mask in observations.masks] masks = torch.tensor(np.array(masks), dtype=torch.float32, device=device) # scale image data (memory reduction) @@ -399,10 +399,14 @@ def fitness_closure() -> torch.Tensor: ).astype(np.uint8) # upscale render current_best_render = cv2.resize( - current_best_render, (images[offset].shape[1], images[offset].shape[0]) + current_best_render, + ( + observations.images[offset].shape[1], + observations.images[offset].shape[0], + ), ) overlay = overlay_mask( - images[offset], + observations.images[offset], current_best_render, scale=1.0, ) diff --git a/cli/rr_hydra.py b/cli/rr_hydra.py index 955e632..76831d5 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -11,7 +11,7 @@ load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, parse_camera_info, - parse_hydra_data, + parse_hydra_observations, ) from roboreg.util import ( clean_xyz, @@ -175,7 +175,7 @@ def main(): joint_states_files = find_files(args.path, args.joint_states_pattern) mask_files = find_files(args.path, args.mask_pattern) depth_files = find_files(args.path, args.depth_pattern) - joint_states, masks, depths = parse_hydra_data( + observations = parse_hydra_observations( joint_states_files=joint_states_files, mask_files=mask_files, depth_files=depth_files, @@ -183,7 +183,7 @@ def main(): height, width, intrinsics = parse_camera_info(args.camera_info_file) # instantiate robot - batch_size = len(joint_states) + batch_size = len(observations.joint_states) if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -201,7 +201,7 @@ def main(): ) mesh_container = TorchMeshContainer( meshes=robot_data.meshes, - batch_size=len(joint_states), + batch_size=len(observations.joint_states), device=device, ) kinematics = TorchKinematics( @@ -217,13 +217,15 @@ def main(): # perform forward kinematics joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) robot.configure(joint_states) # turn depths into xyzs intrinsics = torch.tensor(intrinsics, dtype=torch.float32, device=device) - depths = torch.tensor(np.array(depths), dtype=torch.float32, device=device) + depths = torch.tensor( + np.array(observations.depths), dtype=torch.float32, device=device + ) xyzs = depth_to_xyz( depth=depths, intrinsics=intrinsics, @@ -276,7 +278,7 @@ def main(): dtype=torch.float32, device=device, ) - for xyz, mask in zip(xyzs, masks) + for xyz, mask in zip(xyzs, observations.masks) ] # sample N points per mesh diff --git a/cli/rr_mono_dr.py b/cli/rr_mono_dr.py index 9d6dc73..42fbbef 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -22,7 +22,7 @@ find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, - parse_mono_data, + parse_monocular_observations, ) from roboreg.losses import soft_dice_loss from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask @@ -176,7 +176,7 @@ def main() -> None: image_files = find_files(args.path, args.image_pattern) joint_states_files = find_files(args.path, args.joint_states_pattern) mask_files = find_files(args.path, args.mask_pattern) - images, joint_states, masks = parse_mono_data( + observations = parse_monocular_observations( image_files=image_files, joint_states_files=joint_states_files, mask_files=mask_files, @@ -184,12 +184,12 @@ def main() -> None: # pre-process data joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - targets = [mask_distance_transform(mask) for mask in masks] + targets = [mask_distance_transform(mask) for mask in observations.masks] elif mode == REGISTRATION_MODE.SEGMENTATION: - targets = [mask_exponential_decay(mask) for mask in masks] + targets = [mask_exponential_decay(mask) for mask in observations.masks] else: raise ValueError("Invalid registration mode.") targets = torch.tensor( @@ -298,7 +298,7 @@ def main() -> None: # display optimization progress if args.display_progress: render = renders["camera"][0].squeeze().detach().cpu().numpy() - image = images[0] + image = observations.images[0] render_overlay = overlay_mask( image, (render * 255.0).astype(np.uint8), @@ -307,7 +307,7 @@ def main() -> None: # difference left / right render / mask difference = ( cv2.cvtColor( - np.abs(render - masks[0].astype(np.float32) / 255.0), + np.abs(render - observations.masks[0].astype(np.float32) / 255.0), cv2.COLOR_GRAY2BGR, ) * 255.0 @@ -315,7 +315,7 @@ def main() -> None: # overlay segmentation mask segmentation_overlay = overlay_mask( image, - masks[0], + observations.masks[0], mode="b", scale=1.0, ) @@ -343,8 +343,10 @@ def main() -> None: for i, render in enumerate(renders): render = render.squeeze().cpu().numpy() - overlay = overlay_mask(images[i], (render * 255.0).astype(np.uint8), scale=1.0) - difference = np.abs(render - masks[i].astype(np.float32) / 255.0) + overlay = overlay_mask( + observations.images[i], (render * 255.0).astype(np.uint8), scale=1.0 + ) + difference = np.abs(render - observations.masks[i].astype(np.float32) / 255.0) cv2.imwrite(os.path.join(args.path, f"dr_overlay_{i}.png"), overlay) cv2.imwrite( diff --git a/cli/rr_stereo_dr.py b/cli/rr_stereo_dr.py index bb158ed..a93c940 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -22,7 +22,7 @@ find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, - parse_stereo_data, + parse_stereo_observations, ) from roboreg.losses import soft_dice_loss from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask @@ -208,26 +208,32 @@ def main() -> None: joint_states_files = find_files(args.path, args.joint_states_pattern) left_mask_files = find_files(args.path, args.left_mask_pattern) right_mask_files = find_files(args.path, args.right_mask_pattern) - left_images, right_images, joint_states, left_masks, right_masks = ( - parse_stereo_data( - left_image_files=left_image_files, - right_image_files=right_image_files, - joint_states_files=joint_states_files, - left_mask_files=left_mask_files, - right_mask_files=right_mask_files, - ) + observations = parse_stereo_observations( + left_image_files=left_image_files, + right_image_files=right_image_files, + joint_states_files=joint_states_files, + left_mask_files=left_mask_files, + right_mask_files=right_mask_files, ) # pre-process data joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - left_targets = [mask_distance_transform(mask) for mask in left_masks] - right_targets = [mask_distance_transform(mask) for mask in right_masks] + left_targets = [ + mask_distance_transform(mask) for mask in observations.left_masks + ] + right_targets = [ + mask_distance_transform(mask) for mask in observations.right_masks + ] elif mode == REGISTRATION_MODE.SEGMENTATION: - left_targets = [mask_exponential_decay(mask) for mask in left_masks] - right_targets = [mask_exponential_decay(mask) for mask in right_masks] + left_targets = [ + mask_exponential_decay(mask) for mask in observations.left_masks + ] + right_targets = [ + mask_exponential_decay(mask) for mask in observations.right_masks + ] else: raise ValueError("Invalid registration mode.") left_targets = torch.tensor( @@ -352,7 +358,7 @@ def main() -> None: if args.display_progress: render_overlays = [] left_render = renders["left"][0].squeeze().detach().cpu().numpy() - left_image = left_images[0] + left_image = observations.left_images[0] render_overlays.append( overlay_mask( left_image, @@ -361,7 +367,7 @@ def main() -> None: ) ) right_render = renders["right"][0].squeeze().detach().cpu().numpy() - right_image = right_images[0] + right_image = observations.right_images[0] render_overlays.append( overlay_mask( right_image, @@ -374,7 +380,10 @@ def main() -> None: differences.append( ( cv2.cvtColor( - np.abs(left_render - left_masks[0].astype(np.float32) / 255.0), + np.abs( + left_render + - observations.left_masks[0].astype(np.float32) / 255.0 + ), cv2.COLOR_GRAY2BGR, ) * 255.0 @@ -384,7 +393,8 @@ def main() -> None: ( cv2.cvtColor( np.abs( - right_render - right_masks[0].astype(np.float32) / 255.0 + right_render + - observations.right_masks[0].astype(np.float32) / 255.0 ), cv2.COLOR_GRAY2BGR, ) @@ -396,7 +406,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( left_image, - left_masks[0], + observations.left_masks[0], mode="b", scale=1.0, ) @@ -404,7 +414,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( right_image, - right_masks[0], + observations.right_masks[0], mode="b", scale=1.0, ) @@ -440,14 +450,20 @@ def main() -> None: left_render = left_render.squeeze().cpu().numpy() right_render = right_render.squeeze().cpu().numpy() left_overlay = overlay_mask( - left_images[i], (left_render * 255.0).astype(np.uint8), scale=1.0 + observations.left_images[i], + (left_render * 255.0).astype(np.uint8), + scale=1.0, ) right_overlay = overlay_mask( - right_images[i], (right_render * 255.0).astype(np.uint8), scale=1.0 + observations.right_images[i], + (right_render * 255.0).astype(np.uint8), + scale=1.0, + ) + left_difference = np.abs( + left_render - observations.left_masks[i].astype(np.float32) / 255.0 ) - left_difference = np.abs(left_render - left_masks[i].astype(np.float32) / 255.0) right_difference = np.abs( - right_render - right_masks[i].astype(np.float32) / 255.0 + right_render - observations.right_masks[i].astype(np.float32) / 255.0 ) cv2.imwrite(os.path.join(args.path, f"left_dr_overlay_{i}.png"), left_overlay) diff --git a/roboreg/core/structs.py b/roboreg/core/structs.py index cc95da1..a82f4a5 100644 --- a/roboreg/core/structs.py +++ b/roboreg/core/structs.py @@ -1,12 +1,29 @@ import abc from collections import OrderedDict +from dataclasses import dataclass from pathlib import Path from typing import Dict, List, Optional, Tuple, Union import numpy as np import torch -from roboreg.io import Mesh + +@dataclass +class Mesh: + r"""Dataclass to hold mesh data.""" + + vertices: np.ndarray + faces: np.ndarray + + +@dataclass +class RobotData: + r"""Data needed to construct a Robot.""" + + meshes: Dict[str, Mesh] + urdf: str + root_link_name: str + end_link_name: str class TorchMeshContainer: diff --git a/roboreg/io/meshes.py b/roboreg/io/meshes.py index cfafc2c..27b3c3b 100644 --- a/roboreg/io/meshes.py +++ b/roboreg/io/meshes.py @@ -1,4 +1,3 @@ -from dataclasses import dataclass from pathlib import Path from typing import Dict, Union @@ -6,13 +5,7 @@ import numpy as np import trimesh - -@dataclass -class Mesh: - r"""Dataclass to hold mesh data.""" - - vertices: np.ndarray - faces: np.ndarray +from roboreg.core.structs import Mesh def load_mesh(path: Union[Path, str]) -> Mesh: diff --git a/roboreg/io/parsers.py b/roboreg/io/parsers.py index 50705ab..cfee0e2 100644 --- a/roboreg/io/parsers.py +++ b/roboreg/io/parsers.py @@ -7,6 +7,9 @@ import yaml from pytorch_kinematics import urdf_parser_py +from roboreg.reg.img.request import MonocularObservations, StereoObservations +from roboreg.reg.pcl.request import HydraObservations + class URDFParser: __slots__ = ["_urdf", "_robot"] @@ -312,11 +315,11 @@ def parse_camera_info( return height, width, intrinsic_matrix -def parse_hydra_data( +def parse_hydra_observations( joint_states_files: List[Path], mask_files: List[Path], depth_files: List[Path], -) -> Tuple[List[np.ndarray], List[np.ndarray], List[np.ndarray]]: +) -> HydraObservations: r"""Parse data for Hydra registration. Args: @@ -325,10 +328,7 @@ def parse_hydra_data( depth_files (List[Path]): Depth files. Note that depth values are expected in meters. Returns: - Tuple[List[np.ndarray],List[np.ndarray],List[np.ndarray]]: - - Joint states. - - Masks of shape HxW. - - Point clouds of shape HxWx3. + HydraObservations: Data for Hydra registration. """ if len(joint_states_files) == 0 or len(mask_files) == 0 or len(depth_files) == 0: raise ValueError("No files found.") @@ -348,27 +348,19 @@ def parse_hydra_data( joint_states = [np.load(f) for f in joint_states_files] masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in mask_files] depths = [np.load(f) for f in depth_files] - if not all([mask.dtype == np.uint8 for mask in masks]): - raise ValueError("Masks must be of type np.uint8.") - if not all([np.all(mask >= 0) and np.all(mask <= 255) for mask in masks]): - raise ValueError("Masks must be in the range [0, 255].") if not all( [mask.shape[:2] == depth.shape[:2] for mask, depth in zip(masks, depths)] ): raise ValueError("Mask and depth shapes do not match.") - if not all(mask.ndim == 2 for mask in masks): - raise ValueError("Masks must be 2D.") - if not all(depth.ndim == 2 for depth in depths): - raise ValueError("Depths must be 2D.") - return joint_states, masks, depths + return HydraObservations(joint_states=joint_states, masks=masks, depths=depths) -def parse_mono_data( +def parse_monocular_observations( image_files: List[Path], joint_states_files: List[Path], mask_files: List[Path], -) -> Tuple[List[np.ndarray], List[np.ndarray], List[np.ndarray]]: - r"""Parse monocular data. +) -> MonocularObservations: + r"""Parse monocular observations. Args: image_files (List[Path]): Image files. @@ -376,10 +368,7 @@ def parse_mono_data( mask_files (List[Path]): Mask files. Returns: - Tuple[List[np.ndarray],List[np.ndarray],List[np.ndarray]]: - - Images of shape HxWx3. - - Joint states. - - Masks of shape HxW. + MonocularObservations: Data for monocular registration. """ if len(image_files) != len(joint_states_files) or len(image_files) != len( mask_files @@ -394,37 +383,21 @@ def parse_mono_data( images = [cv2.imread(f) for f in image_files] joint_states = [np.load(f) for f in joint_states_files] masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in mask_files] - if not all([mask.dtype == np.uint8 for mask in masks]): - raise ValueError("Masks must be of type np.uint8.") - if not all([np.all(mask >= 0) and np.all(mask <= 255) for mask in masks]): - raise ValueError("Masks must be in the range [0, 255].") if not all( [mask.shape[:2] == image.shape[:2] for mask, image in zip(masks, images)] ): raise ValueError("Mask and image shapes do not match.") - if not all(mask.ndim == 2 for mask in masks): - raise ValueError("Masks must be 2D.") - if not all(image.ndim == 3 for image in images): - raise ValueError("Images must be 3D.") - if not all(image.shape[-1] == 3 for image in images): - raise ValueError("Images must have 3 channels") - return images, joint_states, masks + return MonocularObservations(images=images, joint_states=joint_states, masks=masks) -def parse_stereo_data( +def parse_stereo_observations( left_image_files: List[Path], right_image_files: List[Path], joint_states_files: List[Path], left_mask_files: List[Path], right_mask_files: List[Path], -) -> Tuple[ - List[np.ndarray], - List[np.ndarray], - List[np.ndarray], - List[np.ndarray], - List[np.ndarray], -]: - r"""Parse stereo data. +) -> StereoObservations: + r"""Parse stereo observations. Args: left_image_files (List[Path]): Left image files. @@ -434,12 +407,7 @@ def parse_stereo_data( right_mask_files (List[Path]): Right mask files. Returns: - Tuple[List[np.ndarray],List[np.ndarray],List[np.ndarray],List[np.ndarray],List[np.ndarray]]: - - Left images of shape HxWx3. - - Right images of shape HxWx3. - - Joint states. - - Left masks of shape HxW. - - Right masks of shape HxW. + StereoObservations: Data for stereo registration. """ if ( len(left_image_files) != len(right_image_files) @@ -463,24 +431,10 @@ def parse_stereo_data( joint_states = [np.load(f) for f in joint_states_files] left_masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in left_mask_files] right_masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in right_mask_files] - if not all([mask.dtype == np.uint8 for mask in left_masks]): - raise ValueError("Left masks must be of type np.uint8.") - if not all([np.all(mask >= 0) and np.all(mask <= 255) for mask in left_masks]): - raise ValueError("Left masks must be in the range [0, 255].") - if not all([mask.dtype == np.uint8 for mask in right_masks]): - raise ValueError("Left masks must be of type np.uint8.") - if not all([np.all(mask >= 0) and np.all(mask <= 255) for mask in right_masks]): - raise ValueError("Left masks must be in the range [0, 255].") - if not all(mask.ndim == 2 for mask in left_masks): - raise ValueError("Left masks must be 2D.") - if not all(image.ndim == 3 for image in left_images): - raise ValueError("Left images must be 3D.") - if not all(image.shape[-1] == 3 for image in left_images): - raise ValueError("Left images must have 3 channels") - if not all(mask.ndim == 2 for mask in right_masks): - raise ValueError("Right masks must be 2D.") - if not all(image.ndim == 3 for image in right_images): - raise ValueError("Right images must be 3D.") - if not all(image.shape[-1] == 3 for image in right_images): - raise ValueError("Right images must have 3 channels") - return left_images, right_images, joint_states, left_masks, right_masks + return StereoObservations( + left_images=left_images, + right_images=right_images, + joint_states=joint_states, + left_masks=left_masks, + right_masks=right_masks, + ) diff --git a/roboreg/io/robot_data.py b/roboreg/io/robot_data.py index 6e49c84..b1fe2ce 100644 --- a/roboreg/io/robot_data.py +++ b/roboreg/io/robot_data.py @@ -1,26 +1,10 @@ -from dataclasses import dataclass from pathlib import Path -from typing import Dict, Union +from typing import Union import rich -from roboreg.io import ( - Mesh, - URDFParser, - apply_mesh_origins, - load_meshes, - simplify_meshes, -) - - -@dataclass -class RobotData: - r"""Data needed to construct a Robot.""" - - meshes: Dict[str, Mesh] - urdf: str - root_link_name: str - end_link_name: str +from roboreg.core.structs import RobotData +from roboreg.io import URDFParser, apply_mesh_origins, load_meshes, simplify_meshes def load_robot_data_from_ros_xacro( diff --git a/roboreg/reg/__init__.py b/roboreg/reg/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/roboreg/reg/_validation.py b/roboreg/reg/_validation.py new file mode 100644 index 0000000..eacee24 --- /dev/null +++ b/roboreg/reg/_validation.py @@ -0,0 +1,43 @@ +from typing import List + +import numpy as np + + +def validate_intrinsics(intrinsics: np.ndarray) -> None: + if intrinsics.shape != (3, 3): + raise ValueError(f"Intrinsics must have shape (3, 3), got {intrinsics.shape}.") + + +def validate_extrinsics(extrinsics: np.ndarray) -> None: + if extrinsics.shape != (4, 4): + raise ValueError( + "Extrinsics must have shape (4, 4), " f"got {extrinsics.shape}." + ) + + +def validate_images(images: List[np.ndarray], name: str) -> None: + for index, image in enumerate(images): + if image.ndim != 3: + raise ValueError(f"{name}[{index}] must be 3D, got shape {image.shape}.") + + if image.shape[-1] != 3: + raise ValueError( + f"{name}[{index}] must have 3 channels, got shape {image.shape}." + ) + + +def validate_masks(masks: List[np.ndarray], name: str) -> None: + for index, mask in enumerate(masks): + if mask.ndim != 2: + raise ValueError(f"{name}[{index}] must be 2D, got shape {mask.shape}.") + + if mask.dtype != np.uint8: + raise ValueError( + f"{name}[{index}] must have dtype np.uint8, got {mask.dtype}." + ) + + +def validate_depths(depths: List[np.ndarray], name: str) -> None: + for index, depth in enumerate(depths): + if depth.ndim != 2: + raise ValueError(f"{name}[{index}] must be 2D, got shape {depth.shape}.") diff --git a/roboreg/reg/img/__init__.py b/roboreg/reg/img/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/roboreg/reg/img/request.py b/roboreg/reg/img/request.py new file mode 100644 index 0000000..fe6d42d --- /dev/null +++ b/roboreg/reg/img/request.py @@ -0,0 +1,96 @@ +from dataclasses import dataclass +from typing import List + +import numpy as np + +from roboreg.core.structs import RobotData +from roboreg.reg._validation import ( + validate_extrinsics, + validate_images, + validate_intrinsics, + validate_masks, +) + + +@dataclass(frozen=True) +class MonocularObservations: + images: List[np.ndarray] + joint_states: List[np.ndarray] + masks: List[np.ndarray] + + def __post_init__(self) -> None: + lengths = { + "images": len(self.images), + "joint_states": len(self.joint_states), + "masks": len(self.masks), + } + + if len(set(lengths.values())) != 1: + raise ValueError( + f"All observation fields must have the same length, got {lengths}." + ) + + if not self.images: + raise ValueError("Expected at least one observation.") + + validate_images(self.images, "images") + validate_masks(self.masks, "masks") + + +@dataclass(frozen=True) +class StereoObservations: + left_images: List[np.ndarray] + right_images: List[np.ndarray] + joint_states: List[np.ndarray] + left_masks: List[np.ndarray] + right_masks: List[np.ndarray] + + def __post_init__(self) -> None: + lengths = { + "left_images": len(self.left_images), + "right_images": len(self.right_images), + "joint_states": len(self.joint_states), + "left_masks": len(self.left_masks), + "right_masks": len(self.right_masks), + } + + if len(set(lengths.values())) != 1: + raise ValueError( + f"All observation fields must have the same length, got {lengths}." + ) + + if not self.left_images: + raise ValueError("Expected at least one observation.") + + validate_images(self.left_images, "left_images") + validate_images(self.right_images, "right_images") + validate_masks(self.left_masks, "left_masks") + validate_masks(self.right_masks, "right_masks") + + +@dataclass(frozen=True) +class MonocularRequest: + intrinsics: np.ndarray + robot_data: RobotData + observations: MonocularObservations + initial_extrinsics: np.ndarray + + def __post_init__(self) -> None: + validate_intrinsics(self.intrinsics) + validate_extrinsics(self.initial_extrinsics) + + +@dataclass(frozen=True) +class StereoRequest: + left_intrinsics: np.ndarray + right_intrinsics: np.ndarray + initial_left_extrinsics: np.ndarray + left_to_right_extrinsics: np.ndarray + robot_data: RobotData + observations: StereoObservations + + def __post_init__(self) -> None: + validate_intrinsics(self.left_intrinsics) + validate_intrinsics(self.right_intrinsics) + validate_extrinsics(self.initial_left_extrinsics) + validate_extrinsics(self.left_to_right_extrinsics) diff --git a/roboreg/reg/pcl/__init__.py b/roboreg/reg/pcl/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/roboreg/reg/pcl/request.py b/roboreg/reg/pcl/request.py new file mode 100644 index 0000000..a6dadec --- /dev/null +++ b/roboreg/reg/pcl/request.py @@ -0,0 +1,49 @@ +from dataclasses import dataclass +from typing import List + +import numpy as np + +from roboreg.core.structs import RobotData +from roboreg.reg._validation import ( + validate_depths, + validate_extrinsics, + validate_intrinsics, + validate_masks, +) + + +@dataclass(frozen=True) +class HydraObservations: + joint_states: List[np.ndarray] + masks: List[np.ndarray] + depths: List[np.ndarray] + + def __post_init__(self): + lengths = { + "joint_states": len(self.joint_states), + "masks": len(self.masks), + "depths": len(self.depths), + } + + if len(set(lengths.values())) != 1: + raise ValueError( + f"All observation fields must have the same length, got {lengths}." + ) + + if not self.joint_states: + raise ValueError("Expected at least one observation.") + + validate_masks(self.masks, "masks") + validate_depths(self.depths, "depths") + + +@dataclass(frozen=True) +class HydraRequest: + intrinsics: np.ndarray + robot_data: RobotData + observations: HydraObservations + initial_extrinsics: np.ndarray + + def __post_init__(self): + validate_intrinsics(self.intrinsics) + validate_extrinsics(self.initial_extrinsics) diff --git a/test/io/test_parsers.py b/test/io/test_parsers.py index b3a13ea..77c8929 100644 --- a/test/io/test_parsers.py +++ b/test/io/test_parsers.py @@ -7,9 +7,9 @@ URDFParser, find_files, parse_camera_info, - parse_hydra_data, - parse_mono_data, - parse_stereo_data, + parse_hydra_observations, + parse_monocular_observations, + parse_stereo_observations, ) @@ -95,97 +95,109 @@ def test_parse_camera_info() -> None: assert intrinsic_matrix.shape == (3, 3), "Intrinsic matrix should be of shape 3x3." -def test_parse_hydra_data() -> None: +def test_parse_hydra_observations() -> None: path = "test/assets/lbr_med7_r800/samples" - joint_states, masks, depths = parse_hydra_data( + observations = parse_hydra_observations( joint_states_files=find_files(path, "joint_states_*.npy"), mask_files=find_files(path, "mask_sam2_left_*.png"), depth_files=find_files(path, "depth_*.npy"), ) assert ( - len(joint_states) == len(masks) == len(depths) + len(observations.joint_states) + == len(observations.masks) + == len(observations.depths) ), "Expected same number of joint states / masks / depths." - assert len(joint_states) >= 1, "Should at least have one sample." - assert masks[0].ndim == 2, "Expected 2D mask." - assert masks[0].dtype == np.uint8, "Expected unsigned integers for mask." - assert np.all(masks[0] >= 0) and np.all( - masks[0] <= 255 + assert len(observations.joint_states) >= 1, "Should at least have one sample." + assert observations.masks[0].ndim == 2, "Expected 2D mask." + assert ( + observations.masks[0].dtype == np.uint8 + ), "Expected unsigned integers for mask." + assert np.all(observations.masks[0] >= 0) and np.all( + observations.masks[0] <= 255 ), "Expected mask in range [0, 255]." - assert depths[0].ndim == 2, "Expected 2D depth map." + assert observations.depths[0].ndim == 2, "Expected 2D depth map." -def test_parse_mono_data() -> None: +def test_parse_monocular_observations() -> None: path = "test/assets/lbr_med7_r800/samples" - images, joint_states, masks = parse_mono_data( + observations = parse_monocular_observations( image_files=find_files(path, "left_image_*.png"), joint_states_files=find_files(path, "joint_states_*.npy"), mask_files=find_files(path, "mask_sam2_left_*.png"), ) assert ( - len(images) == len(joint_states) == len(masks) + len(observations.images) + == len(observations.joint_states) + == len(observations.masks) ), "Expected same number of images / joint states / masks." - assert len(images) >= 1, "Should at least have one sample." - assert images[0].ndim == 3, "Expected 3D image (HxWx3)." - assert images[0].shape[-1] == 3, "Expected 3 color channels." - assert masks[0].ndim == 2, "Expected 2D mask." - assert masks[0].dtype == np.uint8, "Expected unsigned integers for mask." - assert np.all(masks[0] >= 0) and np.all( - masks[0] <= 255 + assert len(observations.images) >= 1, "Should at least have one sample." + assert observations.images[0].ndim == 3, "Expected 3D image (HxWx3)." + assert observations.images[0].shape[-1] == 3, "Expected 3 color channels." + assert observations.masks[0].ndim == 2, "Expected 2D mask." + assert ( + observations.masks[0].dtype == np.uint8 + ), "Expected unsigned integers for mask." + assert np.all(observations.masks[0] >= 0) and np.all( + observations.masks[0] <= 255 ), "Expected mask in range [0, 255]." assert ( - masks[0].shape[:2] == images[0].shape[:2] + observations.masks[0].shape[:2] == observations.images[0].shape[:2] ), "Mask and image dimensions should match." -def test_parse_stereo_data() -> None: +def test_parse_stereo_observations() -> None: path = "test/assets/lbr_med7_r800/samples" - left_images, right_images, joint_states, left_masks, right_masks = ( - parse_stereo_data( - left_image_files=find_files(path, "left_image_*.png"), - right_image_files=find_files(path, "right_image_*.png"), - joint_states_files=find_files(path, "joint_states_*.npy"), - left_mask_files=find_files(path, "mask_sam2_left_*.png"), - right_mask_files=find_files(path, "mask_sam2_right_*.png"), - ) + observations = parse_stereo_observations( + left_image_files=find_files(path, "left_image_*.png"), + right_image_files=find_files(path, "right_image_*.png"), + joint_states_files=find_files(path, "joint_states_*.npy"), + left_mask_files=find_files(path, "mask_sam2_left_*.png"), + right_mask_files=find_files(path, "mask_sam2_right_*.png"), ) assert ( - len(left_images) - == len(right_images) - == len(joint_states) - == len(left_masks) - == len(right_masks) + len(observations.left_images) + == len(observations.right_images) + == len(observations.joint_states) + == len(observations.left_masks) + == len(observations.right_masks) ), "Expected same number of left/right images, joint states, and left/right masks." - assert len(left_images) >= 1, "Should at least have one sample." + assert len(observations.left_images) >= 1, "Should at least have one sample." # Test left data - assert left_images[0].ndim == 3, "Expected 3D left image (HxWx3)." - assert left_images[0].shape[-1] == 3, "Expected 3 color channels for left image." - assert left_masks[0].ndim == 2, "Expected 2D left mask." - assert left_masks[0].dtype == np.uint8, "Expected unsigned integers for left mask." - assert np.all(left_masks[0] >= 0) and np.all( - left_masks[0] <= 255 + assert observations.left_images[0].ndim == 3, "Expected 3D left image (HxWx3)." + assert ( + observations.left_images[0].shape[-1] == 3 + ), "Expected 3 color channels for left image." + assert observations.left_masks[0].ndim == 2, "Expected 2D left mask." + assert ( + observations.left_masks[0].dtype == np.uint8 + ), "Expected unsigned integers for left mask." + assert np.all(observations.left_masks[0] >= 0) and np.all( + observations.left_masks[0] <= 255 ), "Expected left mask in range [0, 255]." # Test right data - assert right_images[0].ndim == 3, "Expected 3D right image (HxWx3)." - assert right_images[0].shape[-1] == 3, "Expected 3 color channels for right image." - assert right_masks[0].ndim == 2, "Expected 2D right mask." + assert observations.right_images[0].ndim == 3, "Expected 3D right image (HxWx3)." + assert ( + observations.right_images[0].shape[-1] == 3 + ), "Expected 3 color channels for right image." + assert observations.right_masks[0].ndim == 2, "Expected 2D right mask." assert ( - right_masks[0].dtype == np.uint8 + observations.right_masks[0].dtype == np.uint8 ), "Expected unsigned integers for right mask." - assert np.all(right_masks[0] >= 0) and np.all( - right_masks[0] <= 255 + assert np.all(observations.right_masks[0] >= 0) and np.all( + observations.right_masks[0] <= 255 ), "Expected right mask in range [0, 255]." # Test dimensions match assert ( - left_masks[0].shape[:2] == left_images[0].shape[:2] + observations.left_masks[0].shape[:2] == observations.left_images[0].shape[:2] ), "Left mask and image dimensions should match." assert ( - right_masks[0].shape[:2] == right_images[0].shape[:2] + observations.right_masks[0].shape[:2] == observations.right_images[0].shape[:2] ), "Right mask and image dimensions should match." @@ -200,6 +212,6 @@ def test_parse_stereo_data() -> None: test_urdf_parser_from_ros_xacro() test_find_files() test_parse_camera_info() - test_parse_hydra_data() - test_parse_mono_data() - test_parse_stereo_data() + test_parse_hydra_observations() + test_parse_monocular_observations() + test_parse_stereo_observations() diff --git a/test/test_hydra_icp.py b/test/test_hydra_icp.py index 97c9cd4..fb99988 100644 --- a/test/test_hydra_icp.py +++ b/test/test_hydra_icp.py @@ -14,11 +14,11 @@ ) from roboreg.io import ( URDFParser, + apply_mesh_origins, find_files, load_meshes, parse_camera_info, - apply_mesh_origins, - parse_hydra_data, + parse_hydra_observations, ) from roboreg.util import ( RegistrationVisualizer, @@ -118,7 +118,7 @@ def test_hydra_icp(): depth_pattern = "depth_*.npy" # load data - joint_states, masks, depths = parse_hydra_data( + observations = parse_hydra_observations( joint_states_files=find_files(path, joint_states_pattern), mask_files=find_files(path, mask_pattern), depth_files=find_files(path, depth_pattern), @@ -139,7 +139,7 @@ def test_hydra_icp(): ) # instantiate mesh - batch_size = len(joint_states) + batch_size = len(observations.joint_states) meshes = TorchMeshContainer( meshes=apply_mesh_origins( meshes=load_meshes( @@ -158,9 +158,9 @@ def test_hydra_icp(): # perform forward kinematics mesh_vertices = meshes.vertices.clone() joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) - ht_lookup = kinematics.mesh_forward_kinematics(joint_states) + ht_lookup = kinematics.forward_kinematics(joint_states) for link_name, ht in ht_lookup.items(): mesh_vertices[ :, @@ -180,7 +180,9 @@ def test_hydra_icp(): # turn depths into xyzs intrinsics = torch.tensor(intrinsics, dtype=torch.float32, device=device) - depths = torch.tensor(np.array(depths), dtype=torch.float32, device=device) + depths = torch.tensor( + np.array(observations.depths), dtype=torch.float32, device=device + ) xyzs = depth_to_xyz(depth=depths, intrinsics=intrinsics, z_max=1.5) # flatten BxHxWx3 -> Bx(H*W)x3 @@ -204,7 +206,7 @@ def test_hydra_icp(): dtype=torch.float32, device=device, ) - for xyz, mask in zip(xyzs, masks) + for xyz, mask in zip(xyzs, observations.masks) ] # sample 5000 points per mesh @@ -249,7 +251,7 @@ def test_hydra_robust_icp() -> None: depth_pattern = "depth_*.npy" # load data - joint_states, masks, depths = parse_hydra_data( + observations = parse_hydra_observations( joint_states_files=find_files(path, joint_states_pattern), mask_files=find_files(path, mask_pattern), depth_files=find_files(path, depth_pattern), @@ -270,7 +272,7 @@ def test_hydra_robust_icp() -> None: ) # instantiate mesh - batch_size = len(joint_states) + batch_size = len(observations.joint_states) meshes = TorchMeshContainer( meshes=apply_mesh_origins( meshes=load_meshes( @@ -289,7 +291,7 @@ def test_hydra_robust_icp() -> None: # perform forward kinematics mesh_vertices = meshes.vertices.clone() joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device + np.array(observations.joint_states), dtype=torch.float32, device=device ) ht_lookup = kinematics.forward_kinematics(joint_states) for link_name, ht in ht_lookup.items(): @@ -310,7 +312,9 @@ def test_hydra_robust_icp() -> None: # turn depths into xyzs intrinsics = torch.tensor(intrinsics, dtype=torch.float32, device=device) - depths = torch.tensor(np.array(depths), dtype=torch.float32, device=device) + depths = torch.tensor( + np.array(observations.depths), dtype=torch.float32, device=device + ) xyzs = depth_to_xyz(depth=depths, intrinsics=intrinsics, z_max=1.5) # flatten BxHxWx3 -> Bx(H*W)x3 @@ -340,7 +344,7 @@ def test_hydra_robust_icp() -> None: dtype=torch.float32, device=device, ) - for xyz, mask in zip(xyzs, masks) + for xyz, mask in zip(xyzs, observations.masks) ] # sample 5000 points per mesh From bd0dccf31fd2496572b86ce3cd8eab666e029704 Mon Sep 17 00:00:00 2001 From: mhubii Date: Fri, 24 Jul 2026 19:05:44 +0100 Subject: [PATCH 02/12] missing typing --- roboreg/reg/pcl/request.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/roboreg/reg/pcl/request.py b/roboreg/reg/pcl/request.py index a6dadec..beaaf25 100644 --- a/roboreg/reg/pcl/request.py +++ b/roboreg/reg/pcl/request.py @@ -18,7 +18,7 @@ class HydraObservations: masks: List[np.ndarray] depths: List[np.ndarray] - def __post_init__(self): + def __post_init__(self) -> None: lengths = { "joint_states": len(self.joint_states), "masks": len(self.masks), @@ -44,6 +44,6 @@ class HydraRequest: observations: HydraObservations initial_extrinsics: np.ndarray - def __post_init__(self): + def __post_init__(self) -> None: validate_intrinsics(self.intrinsics) validate_extrinsics(self.initial_extrinsics) From da68c91313d545b313cd40cd6cf6a34cac7aaec9 Mon Sep 17 00:00:00 2001 From: mhubii Date: Mon, 27 Jul 2026 08:03:30 +0100 Subject: [PATCH 03/12] further refactors --- cli/rr_cam_swarm.py | 26 +-- cli/rr_hydra.py | 151 ++++--------- cli/rr_mono_dr.py | 39 +--- cli/rr_render.py | 18 +- cli/rr_stereo_dr.py | 46 ++-- roboreg/core/robot.py | 36 ++- roboreg/core/structs.py | 10 - roboreg/io/parsers.py | 42 ++-- roboreg/io/robot_data.py | 2 +- roboreg/{reg => registration}/__init__.py | 0 roboreg/{reg => registration}/_validation.py | 12 +- .../img => registration/image}/__init__.py | 0 roboreg/registration/image/config.py | 75 +++++++ .../img => registration/image}/request.py | 24 +- .../point_cloud}/__init__.py | 0 roboreg/registration/point_cloud/config.py | 36 +++ .../point_cloud/hydra.py} | 12 +- .../point_cloud}/request.py | 24 +- roboreg/registration/point_cloud/solver.py | 207 ++++++++++++++++++ roboreg/registration/result.py | 22 ++ test/core/test_robot.py | 18 +- test/core/test_scene.py | 34 +-- test/io/test_parsers.py | 43 ++-- test/test_hydra_icp.py | 36 +-- 24 files changed, 554 insertions(+), 359 deletions(-) rename roboreg/{reg => registration}/__init__.py (100%) rename roboreg/{reg => registration}/_validation.py (80%) rename roboreg/{reg/img => registration/image}/__init__.py (100%) create mode 100644 roboreg/registration/image/config.py rename roboreg/{reg/img => registration/image}/request.py (81%) rename roboreg/{reg/pcl => registration/point_cloud}/__init__.py (100%) create mode 100644 roboreg/registration/point_cloud/config.py rename roboreg/{hydra_icp.py => registration/point_cloud/hydra.py} (97%) rename roboreg/{reg/pcl => registration/point_cloud}/request.py (62%) create mode 100644 roboreg/registration/point_cloud/solver.py create mode 100644 roboreg/registration/result.py diff --git a/cli/rr_cam_swarm.py b/cli/rr_cam_swarm.py index e74fbc7..ae811e0 100644 --- a/cli/rr_cam_swarm.py +++ b/cli/rr_cam_swarm.py @@ -10,8 +10,6 @@ NVDiffRastRenderer, Robot, RobotScene, - TorchKinematics, - TorchMeshContainer, VirtualCamera, ) from roboreg.io import ( @@ -265,18 +263,18 @@ def main() -> None: camera_info_file=args.camera_info_file ) image_files = find_files(args.path, args.image_pattern) - mask_files = find_files(args.path, args.mask_pattern) + target_files = find_files(args.path, args.mask_pattern) joint_states_files = find_files(args.path, args.joint_states_pattern) n_samples = args.n_samples if n_samples > len(image_files): # randomly sample n_samples n_samples = len(image_files) random_indices = np.random.choice(len(image_files), n_samples, replace=False) image_files = np.array(image_files)[random_indices].tolist() - mask_files = np.array(mask_files)[random_indices].tolist() + target_files = np.array(target_files)[random_indices].tolist() joint_states_files = np.array(joint_states_files)[random_indices].tolist() observations = parse_monocular_observations( image_files=image_files, - mask_files=mask_files, + target_files=target_files, joint_states_files=joint_states_files, ) @@ -285,7 +283,7 @@ def main() -> None: np.array(observations.joint_states), dtype=torch.float32, device=device ) n_joint_states = joint_states.shape[0] - masks = [mask_exponential_decay(mask) for mask in observations.masks] + masks = [mask_exponential_decay(mask) for mask in observations.targets] masks = torch.tensor(np.array(masks), dtype=torch.float32, device=device) # scale image data (memory reduction) @@ -345,20 +343,8 @@ def main() -> None: collision=args.collision_meshes, target_reduction=args.target_reduction, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=batch_size, - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=batch_size, device=device ) # instantiate scene diff --git a/cli/rr_hydra.py b/cli/rr_hydra.py index 76831d5..6d27fb0 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -4,8 +4,6 @@ import numpy as np import torch -from roboreg.core import Robot, TorchKinematics, TorchMeshContainer -from roboreg.hydra_icp import hydra_centroid_alignment, hydra_robust_icp from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, @@ -13,15 +11,13 @@ parse_camera_info, parse_hydra_observations, ) -from roboreg.util import ( - clean_xyz, - compute_vertex_normals, - depth_to_xyz, - from_homogeneous, - generate_ht_optical, - mask_extract_extended_boundary, - to_homogeneous, +from roboreg.registration.point_cloud.config import ( + HydraConfig, + HydraRobustICPConfig, + PointCloudConfig, ) +from roboreg.registration.point_cloud.request import HydraRequest +from roboreg.registration.point_cloud.solver import HydraRobustICP from .util.validate import validate_urdf_source @@ -199,116 +195,43 @@ def main(): end_link_name=args.end_link_name, collision=args.collision_meshes, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=len(observations.joint_states), - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, - ) - - # perform forward kinematics - joint_states = torch.tensor( - np.array(observations.joint_states), dtype=torch.float32, device=device - ) - robot.configure(joint_states) - - # turn depths into xyzs - intrinsics = torch.tensor(intrinsics, dtype=torch.float32, device=device) - depths = torch.tensor( - np.array(observations.depths), dtype=torch.float32, device=device - ) - xyzs = depth_to_xyz( - depth=depths, - intrinsics=intrinsics, - z_min=args.z_min, - z_max=args.z_max, - conversion_factor=args.depth_conversion_factor, - ) - - # flatten BxHxWx3 -> Bx(H*W)x3 - xyzs = xyzs.view(-1, height * width, 3) - xyzs = to_homogeneous(xyzs) - ht_optical = generate_ht_optical(xyzs.shape[0], dtype=torch.float32, device=device) - xyzs = torch.matmul(xyzs, ht_optical.transpose(-1, -2)) - xyzs = from_homogeneous(xyzs) - # unflatten - xyzs = xyzs.view(-1, height, width, 3) - xyzs = [xyz.squeeze() for xyz in xyzs.cpu().numpy()] - - # mesh vertices to list - mesh_vertices = from_homogeneous(robot.configured_vertices) - mesh_vertices = [mesh_vertices[i].contiguous() for i in range(batch_size)] - mesh_normals = [] - for i in range(batch_size): - mesh_normals.append( - compute_vertex_normals( - vertices=mesh_vertices[i], faces=robot.mesh_container.faces - ) - ) - - # clean observed vertices and turn into tensor - observed_vertices = [ - torch.tensor( - clean_xyz( - xyz=xyz, - mask=( - mask - if args.no_boundary - else mask_extract_extended_boundary( - mask, - dilation_kernel=np.ones( - [args.dilation_kernel_size, args.dilation_kernel_size] - ), - erosion_kernel=np.ones( - [args.erosion_kernel_size, args.erosion_kernel_size] - ), - ) - ), + # prepare + config = HydraRobustICPConfig( + HydraConfig( + reference_points_per_mesh=args.number_of_points, + observation=PointCloudConfig( + z_min=args.z_min, + z_max=args.z_max, + depth_conversion_factor=args.depth_conversion_factor, + use_mask_boundary=not args.no_boundary, + dilation_kernel_size=args.dilation_kernel_size, + erosion_kernel_size=args.erosion_kernel_size, ), - dtype=torch.float32, - device=device, + max_correspondence_distance=args.max_distance, + ) + ) + hydra_robust_icp = HydraRobustICP(config=config, device=device) + result = hydra_robust_icp( + request=HydraRequest( + intrinsics=intrinsics, + robot_data=robot_data, + observations=observations, ) - for xyz, mask in zip(xyzs, observations.masks) - ] - - # sample N points per mesh - for i in range(batch_size): - idx = torch.randperm(mesh_vertices[i].shape[0])[: args.number_of_points] - mesh_vertices[i] = mesh_vertices[i][idx] - mesh_normals[i] = mesh_normals[i][idx] - - HT_init = hydra_centroid_alignment(observed_vertices, mesh_vertices) - HT = hydra_robust_icp( - HT_init, - observed_vertices, - mesh_vertices, - mesh_normals, - max_distance=args.max_distance, - outer_max_iter=args.outer_max_iter, - inner_max_iter=args.inner_max_iter, ) - # visualize - if args.display_results: - from roboreg.util import RegistrationVisualizer + # TODO update visualization + # # visualize + # if args.display_results: + # from roboreg.util import RegistrationVisualizer - visualizer = RegistrationVisualizer() - visualizer(mesh_vertices=mesh_vertices, observed_vertices=observed_vertices) - visualizer( - mesh_vertices=mesh_vertices, - observed_vertices=observed_vertices, - HT=torch.linalg.inv(HT), - ) + # visualizer = RegistrationVisualizer() + # visualizer(mesh_vertices=mesh_vertices, observed_vertices=observed_vertices) + # visualizer( + # mesh_vertices=mesh_vertices, + # observed_vertices=observed_vertices, + # HT=torch.linalg.inv(HT), + # ) # to numpy HT = HT.cpu().numpy() diff --git a/cli/rr_mono_dr.py b/cli/rr_mono_dr.py index 42fbbef..d9b1882 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -10,14 +10,7 @@ import rich.progress import torch -from roboreg.core import ( - NVDiffRastRenderer, - Robot, - RobotScene, - TorchKinematics, - TorchMeshContainer, - VirtualCamera, -) +from roboreg.core import NVDiffRastRenderer, Robot, RobotScene, VirtualCamera from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, @@ -175,11 +168,11 @@ def main() -> None: # load data image_files = find_files(args.path, args.image_pattern) joint_states_files = find_files(args.path, args.joint_states_pattern) - mask_files = find_files(args.path, args.mask_pattern) + target_files = find_files(args.path, args.mask_pattern) observations = parse_monocular_observations( image_files=image_files, joint_states_files=joint_states_files, - mask_files=mask_files, + target_files=target_files, ) # pre-process data @@ -187,9 +180,9 @@ def main() -> None: np.array(observations.joint_states), dtype=torch.float32, device=device ) if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - targets = [mask_distance_transform(mask) for mask in observations.masks] + targets = [mask_distance_transform(mask) for mask in observations.targets] elif mode == REGISTRATION_MODE.SEGMENTATION: - targets = [mask_exponential_decay(mask) for mask in observations.masks] + targets = [mask_exponential_decay(mask) for mask in observations.targets] else: raise ValueError("Invalid registration mode.") targets = torch.tensor( @@ -220,20 +213,8 @@ def main() -> None: end_link_name=args.end_link_name, collision=args.collision_meshes, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=joint_states.shape[0], - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=joint_states.shape[0], device=device ) # instantiate scene @@ -307,7 +288,7 @@ def main() -> None: # difference left / right render / mask difference = ( cv2.cvtColor( - np.abs(render - observations.masks[0].astype(np.float32) / 255.0), + np.abs(render - observations.targets[0].astype(np.float32) / 255.0), cv2.COLOR_GRAY2BGR, ) * 255.0 @@ -315,7 +296,7 @@ def main() -> None: # overlay segmentation mask segmentation_overlay = overlay_mask( image, - observations.masks[0], + observations.targets[0], mode="b", scale=1.0, ) @@ -346,7 +327,7 @@ def main() -> None: overlay = overlay_mask( observations.images[i], (render * 255.0).astype(np.uint8), scale=1.0 ) - difference = np.abs(render - observations.masks[i].astype(np.float32) / 255.0) + difference = np.abs(render - observations.targets[i].astype(np.float32) / 255.0) cv2.imwrite(os.path.join(args.path, f"dr_overlay_{i}.png"), overlay) cv2.imwrite( diff --git a/cli/rr_render.py b/cli/rr_render.py index 65be955..a72252a 100644 --- a/cli/rr_render.py +++ b/cli/rr_render.py @@ -12,8 +12,6 @@ NVDiffRastRenderer, Robot, RobotScene, - TorchKinematics, - TorchMeshContainer, VirtualCamera, ) from roboreg.io import ( @@ -156,20 +154,8 @@ def main(): end_link_name=args.end_link_name, collision=args.collision_meshes, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=args.batch_size, - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=args.batch_size, device=device ) scene = RobotScene( cameras=camera, diff --git a/cli/rr_stereo_dr.py b/cli/rr_stereo_dr.py index a93c940..30de37f 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -14,8 +14,6 @@ NVDiffRastRenderer, Robot, RobotScene, - TorchKinematics, - TorchMeshContainer, VirtualCamera, ) from roboreg.io import ( @@ -206,14 +204,14 @@ def main() -> None: left_image_files = find_files(args.path, args.left_image_pattern) right_image_files = find_files(args.path, args.right_image_pattern) joint_states_files = find_files(args.path, args.joint_states_pattern) - left_mask_files = find_files(args.path, args.left_mask_pattern) - right_mask_files = find_files(args.path, args.right_mask_pattern) + left_target_files = find_files(args.path, args.left_mask_pattern) + right_target_files = find_files(args.path, args.right_mask_pattern) observations = parse_stereo_observations( left_image_files=left_image_files, right_image_files=right_image_files, joint_states_files=joint_states_files, - left_mask_files=left_mask_files, - right_mask_files=right_mask_files, + left_target_files=left_target_files, + right_target_files=right_target_files, ) # pre-process data @@ -222,17 +220,17 @@ def main() -> None: ) if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: left_targets = [ - mask_distance_transform(mask) for mask in observations.left_masks + mask_distance_transform(mask) for mask in observations.left_targets ] right_targets = [ - mask_distance_transform(mask) for mask in observations.right_masks + mask_distance_transform(mask) for mask in observations.right_targets ] elif mode == REGISTRATION_MODE.SEGMENTATION: left_targets = [ - mask_exponential_decay(mask) for mask in observations.left_masks + mask_exponential_decay(mask) for mask in observations.left_targets ] right_targets = [ - mask_exponential_decay(mask) for mask in observations.right_masks + mask_exponential_decay(mask) for mask in observations.right_targets ] else: raise ValueError("Invalid registration mode.") @@ -274,20 +272,8 @@ def main() -> None: end_link_name=args.end_link_name, collision=args.collision_meshes, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=joint_states.shape[0], - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=joint_states.shape[0], device=device ) # instantiate scene @@ -382,7 +368,7 @@ def main() -> None: cv2.cvtColor( np.abs( left_render - - observations.left_masks[0].astype(np.float32) / 255.0 + - observations.left_targets[0].astype(np.float32) / 255.0 ), cv2.COLOR_GRAY2BGR, ) @@ -394,7 +380,7 @@ def main() -> None: cv2.cvtColor( np.abs( right_render - - observations.right_masks[0].astype(np.float32) / 255.0 + - observations.right_targets[0].astype(np.float32) / 255.0 ), cv2.COLOR_GRAY2BGR, ) @@ -406,7 +392,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( left_image, - observations.left_masks[0], + observations.left_targets[0], mode="b", scale=1.0, ) @@ -414,7 +400,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( right_image, - observations.right_masks[0], + observations.right_targets[0], mode="b", scale=1.0, ) @@ -460,10 +446,10 @@ def main() -> None: scale=1.0, ) left_difference = np.abs( - left_render - observations.left_masks[i].astype(np.float32) / 255.0 + left_render - observations.left_targets[i].astype(np.float32) / 255.0 ) right_difference = np.abs( - right_render - observations.right_masks[i].astype(np.float32) / 255.0 + right_render - observations.right_targets[i].astype(np.float32) / 255.0 ) cv2.imwrite(os.path.join(args.path, f"left_dr_overlay_{i}.png"), left_overlay) diff --git a/roboreg/core/robot.py b/roboreg/core/robot.py index b2ab4ad..f21bccd 100644 --- a/roboreg/core/robot.py +++ b/roboreg/core/robot.py @@ -1,9 +1,20 @@ -from typing import Union +from dataclasses import dataclass +from typing import Dict, Union import torch from .kinematics import TorchKinematics -from .structs import TorchMeshContainer +from .structs import Mesh, TorchMeshContainer + + +@dataclass +class RobotData: + r"""Data needed to construct a Robot.""" + + meshes: Dict[str, Mesh] + urdf: str + root_link_name: str + end_link_name: str class Robot: @@ -23,6 +34,27 @@ def __init__( ) self._device = mesh_container.device + @classmethod + def from_robot_data( + cls, + robot_data: RobotData, + batch_size: int, + device: Union[torch.device, str] = "cuda", + ) -> "Robot": + return Robot( + mesh_container=TorchMeshContainer( + meshes=robot_data.meshes, + batch_size=batch_size, + device=device, + ), + kinematics=TorchKinematics( + urdf=robot_data.urdf, + root_link_name=robot_data.root_link_name, + end_link_name=robot_data.end_link_name, + device=device, + ), + ) + def configure( self, q: torch.FloatTensor, ht_root: torch.FloatTensor = None ) -> None: diff --git a/roboreg/core/structs.py b/roboreg/core/structs.py index a82f4a5..803a375 100644 --- a/roboreg/core/structs.py +++ b/roboreg/core/structs.py @@ -16,16 +16,6 @@ class Mesh: faces: np.ndarray -@dataclass -class RobotData: - r"""Data needed to construct a Robot.""" - - meshes: Dict[str, Mesh] - urdf: str - root_link_name: str - end_link_name: str - - class TorchMeshContainer: r"""Compatability utility structure for NVDiffRast rendering and pytorch-kinematics. diff --git a/roboreg/io/parsers.py b/roboreg/io/parsers.py index cfee0e2..5fe1e12 100644 --- a/roboreg/io/parsers.py +++ b/roboreg/io/parsers.py @@ -7,8 +7,8 @@ import yaml from pytorch_kinematics import urdf_parser_py -from roboreg.reg.img.request import MonocularObservations, StereoObservations -from roboreg.reg.pcl.request import HydraObservations +from roboreg.registration.image.request import MonocularObservations, StereoObservations +from roboreg.registration.point_cloud.request import HydraObservations class URDFParser: @@ -358,44 +358,46 @@ def parse_hydra_observations( def parse_monocular_observations( image_files: List[Path], joint_states_files: List[Path], - mask_files: List[Path], + target_files: List[Path], ) -> MonocularObservations: r"""Parse monocular observations. Args: image_files (List[Path]): Image files. joint_states_files (List[Path]): Joint states files. - mask_files (List[Path]): Mask files. + target_files (List[Path]): Target files. Returns: MonocularObservations: Data for monocular registration. """ if len(image_files) != len(joint_states_files) or len(image_files) != len( - mask_files + target_files ): raise ValueError("Number of images, joint states, masks do not match.") rich.print("Parsing the following files:") rich.print(f"Images: {[f.name for f in image_files]}") rich.print(f"Joint states: {[f.name for f in joint_states_files]}") - rich.print(f"Masks: {[f.name for f in mask_files]}") + rich.print(f"Targets: {[f.name for f in target_files]}") images = [cv2.imread(f) for f in image_files] joint_states = [np.load(f) for f in joint_states_files] - masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in mask_files] + masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in target_files] if not all( [mask.shape[:2] == image.shape[:2] for mask, image in zip(masks, images)] ): raise ValueError("Mask and image shapes do not match.") - return MonocularObservations(images=images, joint_states=joint_states, masks=masks) + return MonocularObservations( + images=images, joint_states=joint_states, targets=masks + ) def parse_stereo_observations( left_image_files: List[Path], right_image_files: List[Path], joint_states_files: List[Path], - left_mask_files: List[Path], - right_mask_files: List[Path], + left_target_files: List[Path], + right_target_files: List[Path], ) -> StereoObservations: r"""Parse stereo observations. @@ -403,8 +405,8 @@ def parse_stereo_observations( left_image_files (List[Path]): Left image files. right_image_files (List[Path]): Right image files. joint_states_files (List[Path]): Joint states files. - left_mask_files (List[Path]): Left mask files. - right_mask_files (List[Path]): Right mask files. + left_target_files (List[Path]): Left target files. + right_target_files (List[Path]): Right target files. Returns: StereoObservations: Data for stereo registration. @@ -412,8 +414,8 @@ def parse_stereo_observations( if ( len(left_image_files) != len(right_image_files) or len(left_image_files) != len(joint_states_files) - or len(left_image_files) != len(left_mask_files) - or len(left_image_files) != len(right_mask_files) + or len(left_image_files) != len(left_target_files) + or len(left_image_files) != len(right_target_files) ): raise ValueError( "Number of left / right images, joint states, left / right masks do not match." @@ -423,18 +425,18 @@ def parse_stereo_observations( rich.print(f"Left images: {[f.name for f in left_image_files]}") rich.print(f"Right images: {[f.name for f in right_image_files]}") rich.print(f"Joint states: {[f.name for f in joint_states_files]}") - rich.print(f"Left masks: {[f.name for f in left_mask_files]}") - rich.print(f"Right masks: {[f.name for f in right_mask_files]}") + rich.print(f"Left targets: {[f.name for f in left_target_files]}") + rich.print(f"Right targets: {[f.name for f in right_target_files]}") left_images = [cv2.imread(f) for f in left_image_files] right_images = [cv2.imread(f) for f in right_image_files] joint_states = [np.load(f) for f in joint_states_files] - left_masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in left_mask_files] - right_masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in right_mask_files] + left_targets = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in left_target_files] + right_targets = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in right_target_files] return StereoObservations( left_images=left_images, right_images=right_images, joint_states=joint_states, - left_masks=left_masks, - right_masks=right_masks, + left_targets=left_targets, + right_targets=right_targets, ) diff --git a/roboreg/io/robot_data.py b/roboreg/io/robot_data.py index b1fe2ce..39e20b5 100644 --- a/roboreg/io/robot_data.py +++ b/roboreg/io/robot_data.py @@ -3,7 +3,7 @@ import rich -from roboreg.core.structs import RobotData +from roboreg.core.robot import RobotData from roboreg.io import URDFParser, apply_mesh_origins, load_meshes, simplify_meshes diff --git a/roboreg/reg/__init__.py b/roboreg/registration/__init__.py similarity index 100% rename from roboreg/reg/__init__.py rename to roboreg/registration/__init__.py diff --git a/roboreg/reg/_validation.py b/roboreg/registration/_validation.py similarity index 80% rename from roboreg/reg/_validation.py rename to roboreg/registration/_validation.py index eacee24..e66ce3d 100644 --- a/roboreg/reg/_validation.py +++ b/roboreg/registration/_validation.py @@ -10,9 +10,7 @@ def validate_intrinsics(intrinsics: np.ndarray) -> None: def validate_extrinsics(extrinsics: np.ndarray) -> None: if extrinsics.shape != (4, 4): - raise ValueError( - "Extrinsics must have shape (4, 4), " f"got {extrinsics.shape}." - ) + raise ValueError(f"Extrinsics must have shape (4, 4), got {extrinsics.shape}.") def validate_images(images: List[np.ndarray], name: str) -> None: @@ -37,7 +35,7 @@ def validate_masks(masks: List[np.ndarray], name: str) -> None: ) -def validate_depths(depths: List[np.ndarray], name: str) -> None: - for index, depth in enumerate(depths): - if depth.ndim != 2: - raise ValueError(f"{name}[{index}] must be 2D, got shape {depth.shape}.") +def validate_targets(targets: List[np.ndarray], name: str) -> None: + for index, target in enumerate(targets): + if target.ndim != 2: + raise ValueError(f"{name}[{index}] must be 2D, got shape {target.shape}.") diff --git a/roboreg/reg/img/__init__.py b/roboreg/registration/image/__init__.py similarity index 100% rename from roboreg/reg/img/__init__.py rename to roboreg/registration/image/__init__.py diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py new file mode 100644 index 0000000..4fbccbf --- /dev/null +++ b/roboreg/registration/image/config.py @@ -0,0 +1,75 @@ +import math +from dataclasses import dataclass, field + + +@dataclass(frozen=True) +class ConvergenceConfig: + max_iterations: int = 400 + tolerance: float = 1.0e-3 + patience: int = 50 + + def __post_init__(self) -> None: + if self.max_iterations <= 0: + raise ValueError("max_iterations must be positive.") + if self.tolerance < 0: + raise ValueError("tolerance must be non-negative.") + if self.patience < 0: + raise ValueError("patience must be non-negative.") + + +@dataclass(frozen=True) +class PlateauSchedulerConfig: + mode: str = "min" + factor: float = 0.1 + patience: int = 50 + threshold: float = 1.0e-4 + + def __post_init__(self) -> None: + if self.factor <= 0 or self.factor >= 1: + raise ValueError("factor must be in the range (0, 1).") + if self.patience < 0: + raise ValueError("patience must be non-negative.") + if self.threshold < 0: + raise ValueError("threshold must be non-negative.") + + +@dataclass(frozen=True) +class DRRegConfig: + optimizer: str = "AdamW" + lr: float = 3.0e-2 + + plateau_scheduler: PlateauSchedulerConfig = field( + default_factory=PlateauSchedulerConfig + ) + convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) + + def __post_init__(self) -> None: + if self.lr <= 0: + raise ValueError("lr must be positive.") + + +@dataclass(frozen=True) +class CSRegConfig: + n_cameras: int = 50 + min_distance: float = 0.5 + max_distance: float = 2.0 + angle_range: float = math.pi + + inertia_weight: float = 0.7 + cognitive_coefficient: float = 1.5 + social_coefficient: float = 1.5 + + convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) + + def __post_init__(self) -> None: + if self.n_cameras <= 0: + raise ValueError("n_cameras must be positive.") + + if self.min_distance <= 0: + raise ValueError("min_distance must be positive.") + + if self.max_distance <= self.min_distance: + raise ValueError("max_distance must be greater than min_distance.") + + if self.angle_range <= 0: + raise ValueError("angle_range must be positive.") diff --git a/roboreg/reg/img/request.py b/roboreg/registration/image/request.py similarity index 81% rename from roboreg/reg/img/request.py rename to roboreg/registration/image/request.py index fe6d42d..832c336 100644 --- a/roboreg/reg/img/request.py +++ b/roboreg/registration/image/request.py @@ -3,12 +3,12 @@ import numpy as np -from roboreg.core.structs import RobotData -from roboreg.reg._validation import ( +from roboreg.core.robot import RobotData +from roboreg.registration._validation import ( validate_extrinsics, validate_images, validate_intrinsics, - validate_masks, + validate_targets, ) @@ -16,13 +16,13 @@ class MonocularObservations: images: List[np.ndarray] joint_states: List[np.ndarray] - masks: List[np.ndarray] + targets: List[np.ndarray] def __post_init__(self) -> None: lengths = { "images": len(self.images), "joint_states": len(self.joint_states), - "masks": len(self.masks), + "targets": len(self.targets), } if len(set(lengths.values())) != 1: @@ -34,7 +34,7 @@ def __post_init__(self) -> None: raise ValueError("Expected at least one observation.") validate_images(self.images, "images") - validate_masks(self.masks, "masks") + validate_targets(self.targets, "targets") @dataclass(frozen=True) @@ -42,16 +42,16 @@ class StereoObservations: left_images: List[np.ndarray] right_images: List[np.ndarray] joint_states: List[np.ndarray] - left_masks: List[np.ndarray] - right_masks: List[np.ndarray] + left_targets: List[np.ndarray] + right_targets: List[np.ndarray] def __post_init__(self) -> None: lengths = { "left_images": len(self.left_images), "right_images": len(self.right_images), "joint_states": len(self.joint_states), - "left_masks": len(self.left_masks), - "right_masks": len(self.right_masks), + "left_targets": len(self.left_targets), + "right_targets": len(self.right_targets), } if len(set(lengths.values())) != 1: @@ -64,8 +64,8 @@ def __post_init__(self) -> None: validate_images(self.left_images, "left_images") validate_images(self.right_images, "right_images") - validate_masks(self.left_masks, "left_masks") - validate_masks(self.right_masks, "right_masks") + validate_targets(self.left_targets, "left_targets") + validate_targets(self.right_targets, "right_targets") @dataclass(frozen=True) diff --git a/roboreg/reg/pcl/__init__.py b/roboreg/registration/point_cloud/__init__.py similarity index 100% rename from roboreg/reg/pcl/__init__.py rename to roboreg/registration/point_cloud/__init__.py diff --git a/roboreg/registration/point_cloud/config.py b/roboreg/registration/point_cloud/config.py new file mode 100644 index 0000000..c2d9fa2 --- /dev/null +++ b/roboreg/registration/point_cloud/config.py @@ -0,0 +1,36 @@ +from dataclasses import dataclass, field + + +@dataclass(frozen=True) +class PointCloudConfig: + z_min: float = 0.01 + z_max: float = 2.0 + + depth_conversion_factor: float = 1.0 + + use_mask_boundary: bool = True + dilation_kernel_size: int = 3 + erosion_kernel_size: int = 10 + + +@dataclass(frozen=True) +class HydraConfig: + reference_points_per_mesh: int = 5000 + + observation: PointCloudConfig = field(default_factory=PointCloudConfig) + + max_correspondence_distance: float = 0.1 + rmse_change_tolerance: float = 1e-6 ## convergence ... + + +@dataclass(frozen=True) +class HydraICPConfig: + hydra: HydraConfig = field(default_factory=HydraConfig) + max_iterations: int = 100 + + +@dataclass(frozen=True) +class HydraRobustICPConfig: + hydra: HydraConfig = field(default_factory=HydraConfig) + outer_max_iterations: int = 50 + inner_max_iterations: int = 10 diff --git a/roboreg/hydra_icp.py b/roboreg/registration/point_cloud/hydra.py similarity index 97% rename from roboreg/hydra_icp.py rename to roboreg/registration/point_cloud/hydra.py index 616d6e9..4654732 100644 --- a/roboreg/hydra_icp.py +++ b/roboreg/registration/point_cloud/hydra.py @@ -45,7 +45,7 @@ def kabsch_register( return R, t -def hydra_correspondence_indices( +def correspondence_indices( input: torch.Tensor, target: torch.Tensor, max_distance: float = 0.1 ) -> Tuple[torch.Tensor, torch.Tensor]: r"""For each point in input, find nearest neighbor index in target. @@ -70,7 +70,7 @@ def hydra_correspondence_indices( return matchindices, mask -def hydra_centroid_alignment( +def centroid_alignment( Xs: List[torch.Tensor], Ys: List[torch.Tensor], ) -> torch.Tensor: @@ -101,7 +101,7 @@ def hydra_centroid_alignment( return HT -def hydra_icp( +def point_to_point_icp( HT_init: torch.Tensor, observations: List[torch.Tensor], meshes: List[torch.Tensor], @@ -132,7 +132,7 @@ def hydra_icp( for i in range(len(meshes)): # search correspondences observations_tf = observations[i] @ HT[:3, :3].T + HT[:3, 3] - matchindices, mask = hydra_correspondence_indices( + matchindices, mask = correspondence_indices( observations_tf, meshes[i], max_distance ) @@ -177,7 +177,7 @@ def hydra_icp( return HT -def hydra_robust_icp( +def point_to_plane_robust_icp( HT_init: torch.Tensor, observations: List[torch.Tensor], meshes: List[torch.Tensor], @@ -239,7 +239,7 @@ def hydra_robust_icp( raise ValueError("Length of observations and meshes must be the same.") # search correspondences observations_tf = observations[i] @ HT[:3, :3].T + HT[:3, 3] - matchindices, mask = hydra_correspondence_indices( + matchindices, mask = correspondence_indices( observations_tf, meshes[i], max_distance ) diff --git a/roboreg/reg/pcl/request.py b/roboreg/registration/point_cloud/request.py similarity index 62% rename from roboreg/reg/pcl/request.py rename to roboreg/registration/point_cloud/request.py index beaaf25..09966e7 100644 --- a/roboreg/reg/pcl/request.py +++ b/roboreg/registration/point_cloud/request.py @@ -1,14 +1,13 @@ from dataclasses import dataclass -from typing import List +from typing import List, Tuple import numpy as np -from roboreg.core.structs import RobotData -from roboreg.reg._validation import ( - validate_depths, - validate_extrinsics, +from roboreg.core.robot import RobotData +from roboreg.registration._validation import ( validate_intrinsics, validate_masks, + validate_targets, ) @@ -34,7 +33,18 @@ def __post_init__(self) -> None: raise ValueError("Expected at least one observation.") validate_masks(self.masks, "masks") - validate_depths(self.depths, "depths") + validate_targets(self.depths, "depths") + + for i, (mask, depth) in enumerate(zip(self.masks, self.depths)): + if mask.shape != depth.shape: + raise ValueError( + f"masks[{i}] and depths[{i}] have incompatible shapes: " + f"{mask.shape} and {depth.shape}." + ) + + @property + def shape(self) -> Tuple[int, int]: + return self.depths[0].shape @dataclass(frozen=True) @@ -42,8 +52,6 @@ class HydraRequest: intrinsics: np.ndarray robot_data: RobotData observations: HydraObservations - initial_extrinsics: np.ndarray def __post_init__(self) -> None: validate_intrinsics(self.intrinsics) - validate_extrinsics(self.initial_extrinsics) diff --git a/roboreg/registration/point_cloud/solver.py b/roboreg/registration/point_cloud/solver.py new file mode 100644 index 0000000..482d26a --- /dev/null +++ b/roboreg/registration/point_cloud/solver.py @@ -0,0 +1,207 @@ +from dataclasses import dataclass +from typing import List, Optional + +import numpy as np +import torch + +from roboreg.core.robot import Robot +from roboreg.registration.result import RegistrationResult +from roboreg.util.mask import mask_extract_extended_boundary +from roboreg.util.points import ( + clean_xyz, + compute_vertex_normals, + from_homogeneous, + to_homogeneous, +) +from roboreg.util.transform import depth_to_xyz, generate_ht_optical + +from .config import HydraConfig, HydraICPConfig, HydraRobustICPConfig +from .hydra import centroid_alignment, point_to_plane_robust_icp, point_to_point_icp +from .request import HydraRequest + + +@dataclass(frozen=True) +class _HydraProblem: + observed_vertices: List[torch.Tensor] + reference_vertices: List[torch.Tensor] + reference_normals: Optional[List[torch.Tensor]] = None + + +def _prepare_hydra_problem( + request: HydraRequest, + config: HydraConfig, + device: torch.device, + compute_normals: bool = True, +) -> _HydraProblem: + # 1) construct robot on request + robot = Robot.from_robot_data( + robot_data=request.robot_data, + batch_size=len(request.observations.joint_states), + device=device, + ) + + # 2) to tensor + joint_states = torch.tensor( + np.stack(request.observations.joint_states), dtype=torch.float32, device=device + ) + intrinsics = torch.tensor(request.intrinsics, dtype=torch.float32, device=device) + depths = torch.tensor( + np.stack(request.observations.depths), dtype=torch.float32, device=device + ) + + # 3) perform forward kinematics + robot.configure(joint_states) + + # 4) process depths + xyzs = depth_to_xyz( + depth=depths, + intrinsics=intrinsics, + z_min=config.observation.z_min, + z_max=config.observation.z_max, + conversion_factor=config.observation.depth_conversion_factor, + ) + height, width = request.observations.shape + xyzs = xyzs.view(-1, height * width, 3) # flatten BxHxWx3 -> Bx(H*W)x3 + xyzs = to_homogeneous(xyzs) + ht_optical = generate_ht_optical(xyzs.shape[0], dtype=torch.float32, device=device) + xyzs = torch.matmul(xyzs, ht_optical.transpose(-1, -2)) + xyzs = from_homogeneous(xyzs) + xyzs = xyzs.view(-1, height, width, 3) + xyzs = [xyz.squeeze() for xyz in xyzs.cpu().numpy()] + + # 5) clean observed vertices and turn into tensor + observed_vertices = [ + torch.tensor( + clean_xyz( + xyz=xyz, + mask=( + mask_extract_extended_boundary( + mask, + dilation_kernel=np.ones( + [ + config.observation.dilation_kernel_size, + config.observation.dilation_kernel_size, + ] + ), + erosion_kernel=np.ones( + [ + config.observation.erosion_kernel_size, + config.observation.erosion_kernel_size, + ] + ), + ) + if config.observation.use_mask_boundary + else mask + ), + ), + dtype=torch.float32, + device=device, + ) + for xyz, mask in zip(xyzs, request.observations.masks) + ] + + # mesh vertices to list + batch_size = len(request.observations.joint_states) + + mesh_vertices = from_homogeneous(robot.configured_vertices) + mesh_vertices = [mesh_vertices[i].contiguous() for i in range(batch_size)] + + mesh_normals: list[torch.Tensor] | None = None + + if compute_normals: + mesh_normals = [ + compute_vertex_normals( + vertices=mesh_vertices[i], + faces=robot.mesh_container.faces, + ) + for i in range(batch_size) + ] + + # sample N points per mesh + for i in range(batch_size): + n_points = min( + config.reference_points_per_mesh, + mesh_vertices[i].shape[0], + ) + + idx = torch.randperm( + mesh_vertices[i].shape[0], + device=mesh_vertices[i].device, + )[:n_points] + + mesh_vertices[i] = mesh_vertices[i][idx] + + if mesh_normals is not None: + mesh_normals[i] = mesh_normals[i][idx] + + return _HydraProblem( + observed_vertices=observed_vertices, + reference_vertices=mesh_vertices, + reference_normals=mesh_normals, + ) + + +class HydraICP: + def __init__( + self, + config: HydraICPConfig | None = None, + device: torch.device | str = "cuda", + ) -> None: + self._config = config or HydraICPConfig() + self._device = torch.device(device) + + def __call__(self, request: HydraRequest) -> RegistrationResult: + hydra_problem = _prepare_hydra_problem( + request=request, + config=self._config.hydra, + device=self._device, + compute_normals=False, + ) + HT_init = centroid_alignment( + hydra_problem.observed_vertices, hydra_problem.reference_vertices + ) + HT = point_to_point_icp( + HT_init, + hydra_problem.observed_vertices, + hydra_problem.reference_vertices, + max_distance=self._config.hydra.max_correspondence_distance, + max_iter=self._config.max_iterations, + rmse_change=self._config.hydra.rmse_change_tolerance, + ) + return RegistrationResult( + extrinsics=HT, + ) + + +class HydraRobustICP: + def __init__( + self, + config: HydraRobustICPConfig | None = None, + device: torch.device | str = "cuda", + ) -> None: + self._config = config or HydraRobustICPConfig() + self._device = torch.device(device) + + def __call__(self, request: HydraRequest) -> RegistrationResult: + hydra_problem = _prepare_hydra_problem( + request=request, + config=self._config.hydra, + device=self._device, + compute_normals=True, + ) + HT_init = centroid_alignment( + hydra_problem.observed_vertices, hydra_problem.reference_vertices + ) + HT = point_to_plane_robust_icp( + HT_init, + hydra_problem.observed_vertices, + hydra_problem.reference_vertices, + hydra_problem.reference_normals, + max_distance=self._config.hydra.max_correspondence_distance, + outer_max_iter=self._config.outer_max_iterations, + inner_max_iter=self._config.inner_max_iterations, + rmse_change=self._config.hydra.rmse_change_tolerance, + ) + return RegistrationResult( + extrinsics=HT, + ) diff --git a/roboreg/registration/result.py b/roboreg/registration/result.py new file mode 100644 index 0000000..f3b27df --- /dev/null +++ b/roboreg/registration/result.py @@ -0,0 +1,22 @@ +from dataclasses import dataclass +from enum import Enum + +import torch + + +class TerminationReason(str, Enum): + CONVERGED = "converged" + MAX_ITERATIONS = "max_iterations" + FAILED = "failed" + + +@dataclass +class RegistrationResult: + extrinsics: torch.Tensor + iterations: int + termination_reason: TerminationReason + message: str | None = None + + @property + def converged(self) -> bool: + return self.termination_reason == TerminationReason.CONVERGED diff --git a/test/core/test_robot.py b/test/core/test_robot.py index 4a4c4e0..4fa3c9d 100644 --- a/test/core/test_robot.py +++ b/test/core/test_robot.py @@ -1,6 +1,6 @@ import torch -from roboreg.core import Robot, TorchKinematics, TorchMeshContainer +from roboreg.core import Robot from roboreg.io import load_robot_data_from_urdf_file @@ -12,20 +12,8 @@ def test_robot() -> None: collision=True, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=batch_size, - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=batch_size, device=device ) assert robot.device == torch.device(device), "Robot device mismatch." diff --git a/test/core/test_scene.py b/test/core/test_scene.py index 88e3b0a..6a86cab 100644 --- a/test/core/test_scene.py +++ b/test/core/test_scene.py @@ -10,8 +10,6 @@ NVDiffRastRenderer, Robot, RobotScene, - TorchKinematics, - TorchMeshContainer, VirtualCamera, ) from roboreg.io import find_files, load_robot_data_from_urdf_file @@ -103,20 +101,8 @@ def __init__( root_link_name=root_link_name, end_link_name=end_link_name, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=self.joint_states.shape[0], - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=self.joint_states.shape[0], device=device ) # instantiate scene @@ -238,20 +224,8 @@ def test_single_camera_multiple_poses() -> None: root_link_name="lbr_link_0", end_link_name="lbr_link_7", ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=batch_size, - device=device, - ) - kinematics = TorchKinematics( - urdf=robot_data.urdf, - root_link_name=robot_data.root_link_name, - end_link_name=robot_data.end_link_name, - device=device, - ) - robot = Robot( - mesh_container=mesh_container, - kinematics=kinematics, + robot = Robot.from_robot_data( + robot_data=robot_data, batch_size=batch_size, device=device ) # instantiate scene diff --git a/test/io/test_parsers.py b/test/io/test_parsers.py index 77c8929..f813fec 100644 --- a/test/io/test_parsers.py +++ b/test/io/test_parsers.py @@ -124,26 +124,26 @@ def test_parse_monocular_observations() -> None: observations = parse_monocular_observations( image_files=find_files(path, "left_image_*.png"), joint_states_files=find_files(path, "joint_states_*.npy"), - mask_files=find_files(path, "mask_sam2_left_*.png"), + target_files=find_files(path, "mask_sam2_left_*.png"), ) assert ( len(observations.images) == len(observations.joint_states) - == len(observations.masks) + == len(observations.targets) ), "Expected same number of images / joint states / masks." assert len(observations.images) >= 1, "Should at least have one sample." assert observations.images[0].ndim == 3, "Expected 3D image (HxWx3)." assert observations.images[0].shape[-1] == 3, "Expected 3 color channels." - assert observations.masks[0].ndim == 2, "Expected 2D mask." + assert observations.targets[0].ndim == 2, "Expected 2D mask." assert ( - observations.masks[0].dtype == np.uint8 + observations.targets[0].dtype == np.uint8 ), "Expected unsigned integers for mask." - assert np.all(observations.masks[0] >= 0) and np.all( - observations.masks[0] <= 255 + assert np.all(observations.targets[0] >= 0) and np.all( + observations.targets[0] <= 255 ), "Expected mask in range [0, 255]." assert ( - observations.masks[0].shape[:2] == observations.images[0].shape[:2] + observations.targets[0].shape[:2] == observations.images[0].shape[:2] ), "Mask and image dimensions should match." @@ -153,16 +153,16 @@ def test_parse_stereo_observations() -> None: left_image_files=find_files(path, "left_image_*.png"), right_image_files=find_files(path, "right_image_*.png"), joint_states_files=find_files(path, "joint_states_*.npy"), - left_mask_files=find_files(path, "mask_sam2_left_*.png"), - right_mask_files=find_files(path, "mask_sam2_right_*.png"), + left_target_files=find_files(path, "mask_sam2_left_*.png"), + right_target_files=find_files(path, "mask_sam2_right_*.png"), ) assert ( len(observations.left_images) == len(observations.right_images) == len(observations.joint_states) - == len(observations.left_masks) - == len(observations.right_masks) + == len(observations.left_targets) + == len(observations.right_targets) ), "Expected same number of left/right images, joint states, and left/right masks." assert len(observations.left_images) >= 1, "Should at least have one sample." @@ -171,12 +171,12 @@ def test_parse_stereo_observations() -> None: assert ( observations.left_images[0].shape[-1] == 3 ), "Expected 3 color channels for left image." - assert observations.left_masks[0].ndim == 2, "Expected 2D left mask." + assert observations.left_targets[0].ndim == 2, "Expected 2D left mask." assert ( - observations.left_masks[0].dtype == np.uint8 + observations.left_targets[0].dtype == np.uint8 ), "Expected unsigned integers for left mask." - assert np.all(observations.left_masks[0] >= 0) and np.all( - observations.left_masks[0] <= 255 + assert np.all(observations.left_targets[0] >= 0) and np.all( + observations.left_targets[0] <= 255 ), "Expected left mask in range [0, 255]." # Test right data @@ -184,20 +184,21 @@ def test_parse_stereo_observations() -> None: assert ( observations.right_images[0].shape[-1] == 3 ), "Expected 3 color channels for right image." - assert observations.right_masks[0].ndim == 2, "Expected 2D right mask." + assert observations.right_targets[0].ndim == 2, "Expected 2D right mask." assert ( - observations.right_masks[0].dtype == np.uint8 + observations.right_targets[0].dtype == np.uint8 ), "Expected unsigned integers for right mask." - assert np.all(observations.right_masks[0] >= 0) and np.all( - observations.right_masks[0] <= 255 + assert np.all(observations.right_targets[0] >= 0) and np.all( + observations.right_targets[0] <= 255 ), "Expected right mask in range [0, 255]." # Test dimensions match assert ( - observations.left_masks[0].shape[:2] == observations.left_images[0].shape[:2] + observations.left_targets[0].shape[:2] == observations.left_images[0].shape[:2] ), "Left mask and image dimensions should match." assert ( - observations.right_masks[0].shape[:2] == observations.right_images[0].shape[:2] + observations.right_targets[0].shape[:2] + == observations.right_images[0].shape[:2] ), "Right mask and image dimensions should match." diff --git a/test/test_hydra_icp.py b/test/test_hydra_icp.py index fb99988..5e1ee86 100644 --- a/test/test_hydra_icp.py +++ b/test/test_hydra_icp.py @@ -6,12 +6,6 @@ import transformations as tf from roboreg.core import TorchKinematics, TorchMeshContainer -from roboreg.hydra_icp import ( - hydra_centroid_alignment, - hydra_correspondence_indices, - hydra_icp, - hydra_robust_icp, -) from roboreg.io import ( URDFParser, apply_mesh_origins, @@ -20,6 +14,12 @@ parse_camera_info, parse_hydra_observations, ) +from roboreg.registration.point_cloud.hydra import ( + centroid_alignment, + correspondence_indices, + point_to_point_icp, + point_to_plane_robust_icp, +) from roboreg.util import ( RegistrationVisualizer, clean_xyz, @@ -48,7 +48,7 @@ def test_hydra_centroid_alignment(): for mesh_centroid in mesh_centroids ] - HT = hydra_centroid_alignment(mesh_centroids, observed_centroids) + HT = centroid_alignment(mesh_centroids, observed_centroids) assert torch.allclose(HT, HT_random) @@ -78,7 +78,7 @@ def test_index_shape( # single input input = torch.rand(M, dim) target = torch.rand(N, dim) # e.g. the mesh vertices - matchindices, mask = hydra_correspondence_indices( + matchindices, mask = correspondence_indices( input, target, max_distance=np.sqrt(dim) / 2.0 # remove some elements randomly ) test_index_shape(matchindices, mask, torch.Size([M]), N) @@ -87,7 +87,7 @@ def test_index_shape( batch_size = 2 input = torch.rand(batch_size, M, dim) target = torch.rand(batch_size, N, dim) - matchindices, mask = hydra_correspondence_indices( + matchindices, mask = correspondence_indices( input, target, max_distance=np.sqrt(dim) / 2.0 ) test_index_shape(matchindices, mask, torch.Size([batch_size, M]), N) @@ -98,14 +98,14 @@ def test_index_shape( input = torch.rand(M, dim) target = torch.rand(N, dim) - matchindices, mask = hydra_correspondence_indices( + matchindices, mask = correspondence_indices( input, target, max_distance=np.sqrt(dim) / 2.0 ) test_index_shape(matchindices, mask, torch.Size([M]), N) @pytest.mark.skip(reason="To be fixed.") -def test_hydra_icp(): +def test_hydra_point_to_point_icp(): device = "cuda" if torch.cuda.is_available() else "cpu" ros_package = "lbr_description" xacro_path = "urdf/med7/med7.xacro" @@ -214,8 +214,8 @@ def test_hydra_icp(): idx = torch.randperm(mesh_vertices[i].shape[0])[:5000] mesh_vertices[i] = mesh_vertices[i][idx] - HT_init = hydra_centroid_alignment(observed_vertices, mesh_vertices) - HT = hydra_icp( + HT_init = centroid_alignment(observed_vertices, mesh_vertices) + HT = point_to_point_icp( HT_init, observed_vertices, mesh_vertices, @@ -238,7 +238,7 @@ def test_hydra_icp(): @pytest.mark.skip(reason="To be fixed.") -def test_hydra_robust_icp() -> None: +def test_hydra_point_to_plane_robust_icp() -> None: device = "cuda" if torch.cuda.is_available() else "cpu" ros_package = "lbr_description" xacro_path = "urdf/med7/med7.xacro" @@ -353,8 +353,8 @@ def test_hydra_robust_icp() -> None: mesh_vertices[i] = mesh_vertices[i][idx] mesh_normals[i] = mesh_normals[i][idx] - HT_init = hydra_centroid_alignment(observed_vertices, mesh_vertices) - HT = hydra_robust_icp( + HT_init = centroid_alignment(observed_vertices, mesh_vertices) + HT = point_to_plane_robust_icp( HT_init, observed_vertices, mesh_vertices, @@ -387,5 +387,5 @@ def test_hydra_robust_icp() -> None: # test_hydra_centroid_alignment() # test_hydra_correspondence_indices() - # test_hydra_icp() - test_hydra_robust_icp() + # test_hydra_point_to_point_icp() + test_hydra_point_to_plane_robust_icp() From e983b463aa98c49e1cc6a6069b40cbf6cfa14933 Mon Sep 17 00:00:00 2001 From: mhubii Date: Mon, 27 Jul 2026 21:46:35 +0100 Subject: [PATCH 04/12] finished the hydra refactor --- cli/rr_hydra.py | 54 ++-- roboreg/registration/point_cloud/config.py | 13 +- roboreg/registration/point_cloud/hydra.py | 293 ++++++++++++--------- roboreg/registration/point_cloud/solver.py | 64 +++-- test/test_hydra_icp.py | 110 ++++---- 5 files changed, 299 insertions(+), 235 deletions(-) diff --git a/cli/rr_hydra.py b/cli/rr_hydra.py index 6d27fb0..7c93285 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -12,12 +12,13 @@ parse_hydra_observations, ) from roboreg.registration.point_cloud.config import ( + DepthToPointCloudConfig, HydraConfig, HydraRobustICPConfig, - PointCloudConfig, ) from roboreg.registration.point_cloud.request import HydraRequest -from roboreg.registration.point_cloud.solver import HydraRobustICP +from roboreg.registration.point_cloud.solver import HydraProblem, HydraRobustICP +from roboreg.registration.result import RegistrationResult from .util.validate import validate_urdf_source @@ -163,6 +164,26 @@ def args_factory() -> argparse.Namespace: return parser.parse_args() +def visualize_hydra_result( + problem: HydraProblem, + result: RegistrationResult, +) -> None: + from roboreg.util import RegistrationVisualizer + + visualizer = RegistrationVisualizer() + + visualizer( + mesh_vertices=problem.reference_vertices, + observed_vertices=problem.observed_vertices, + ) + + visualizer( + mesh_vertices=problem.reference_vertices, + observed_vertices=problem.observed_vertices, + HT=torch.linalg.inv(result.extrinsics), + ) + + def main(): args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" @@ -176,10 +197,9 @@ def main(): mask_files=mask_files, depth_files=depth_files, ) - height, width, intrinsics = parse_camera_info(args.camera_info_file) + _, _, intrinsics = parse_camera_info(args.camera_info_file) # instantiate robot - batch_size = len(observations.joint_states) if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -196,11 +216,11 @@ def main(): collision=args.collision_meshes, ) - # prepare + # register config = HydraRobustICPConfig( HydraConfig( reference_points_per_mesh=args.number_of_points, - observation=PointCloudConfig( + depth_to_point_cloud=DepthToPointCloudConfig( z_min=args.z_min, z_max=args.z_max, depth_conversion_factor=args.depth_conversion_factor, @@ -211,7 +231,11 @@ def main(): max_correspondence_distance=args.max_distance, ) ) - hydra_robust_icp = HydraRobustICP(config=config, device=device) + hydra_robust_icp = HydraRobustICP( + config=config, + device=device, + callback=visualize_hydra_result if args.display_results else None, + ) result = hydra_robust_icp( request=HydraRequest( intrinsics=intrinsics, @@ -220,22 +244,8 @@ def main(): ) ) - # TODO update visualization - # # visualize - # if args.display_results: - # from roboreg.util import RegistrationVisualizer - - # visualizer = RegistrationVisualizer() - # visualizer(mesh_vertices=mesh_vertices, observed_vertices=observed_vertices) - # visualizer( - # mesh_vertices=mesh_vertices, - # observed_vertices=observed_vertices, - # HT=torch.linalg.inv(HT), - # ) - # to numpy - HT = HT.cpu().numpy() - np.save(os.path.join(args.path, args.output_file), HT) + np.save(os.path.join(args.path, args.output_file), result.extrinsics.cpu().numpy()) if __name__ == "__main__": diff --git a/roboreg/registration/point_cloud/config.py b/roboreg/registration/point_cloud/config.py index c2d9fa2..79122a0 100644 --- a/roboreg/registration/point_cloud/config.py +++ b/roboreg/registration/point_cloud/config.py @@ -2,10 +2,9 @@ @dataclass(frozen=True) -class PointCloudConfig: +class DepthToPointCloudConfig: z_min: float = 0.01 z_max: float = 2.0 - depth_conversion_factor: float = 1.0 use_mask_boundary: bool = True @@ -17,10 +16,12 @@ class PointCloudConfig: class HydraConfig: reference_points_per_mesh: int = 5000 - observation: PointCloudConfig = field(default_factory=PointCloudConfig) + depth_to_point_cloud: DepthToPointCloudConfig = field( + default_factory=DepthToPointCloudConfig + ) max_correspondence_distance: float = 0.1 - rmse_change_tolerance: float = 1e-6 ## convergence ... + rmse_change_tolerance: float = 1e-6 @dataclass(frozen=True) @@ -32,5 +33,5 @@ class HydraICPConfig: @dataclass(frozen=True) class HydraRobustICPConfig: hydra: HydraConfig = field(default_factory=HydraConfig) - outer_max_iterations: int = 50 - inner_max_iterations: int = 10 + max_outer_iterations: int = 50 + max_inner_iterations: int = 10 diff --git a/roboreg/registration/point_cloud/hydra.py b/roboreg/registration/point_cloud/hydra.py index 4654732..c2a4d71 100644 --- a/roboreg/registration/point_cloud/hydra.py +++ b/roboreg/registration/point_cloud/hydra.py @@ -1,19 +1,20 @@ from typing import List, Tuple import torch -from rich import print from rich.progress import track +from roboreg.registration.result import RegistrationResult, TerminationReason + def kabsch_register( - input: torch.Tensor, target: torch.Tensor + observed_vertices: torch.Tensor, reference_vertices: torch.Tensor ) -> Tuple[torch.Tensor, torch.Tensor]: r"""Kabsch algorithm: https://en.wikipedia.org/wiki/Kabsch_algorithm. - Computes rotation and translation such that input @ R + t = target. + Computes rotation and translation such that observed_vertices @ R + t = reference_vertices. Args: - input (torch.Tensor): input of shape (..., M, 3). - target(torch.Tensor): target of shape (..., M, 3). + observed_vertices (torch.Tensor): Observed vertices of shape (..., M, 3). + reference_vertices (torch.Tensor): Reference vertices of shape (..., M, 3). Returns: Tuple[torch.Tensor,torch.Tensor]: @@ -21,15 +22,15 @@ def kabsch_register( - Translation vector of shape (..., 3). """ # compute centroids - input_centroid = torch.mean(input, dim=-2) - target_centroid = torch.mean(target, dim=-2) + observed_centroid = torch.mean(observed_vertices, dim=-2) + reference_centroid = torch.mean(reference_vertices, dim=-2) # compute centered points - input_centered = input - input_centroid - target_centered = target - target_centroid + observed_centered = observed_vertices - observed_centroid + reference_centered = reference_vertices - reference_centroid # compute covariance matrix - H = target_centered.transpose(-1, -2) @ input_centered + H = reference_centered.transpose(-1, -2) @ observed_centered # compute SVD U, _, V = torch.svd(H) @@ -41,56 +42,60 @@ def kabsch_register( R = V @ E @ U.transpose(-1, -2) # compute translation - t = target_centroid - input_centroid @ R + t = reference_centroid - observed_centroid @ R return R, t def correspondence_indices( - input: torch.Tensor, target: torch.Tensor, max_distance: float = 0.1 + observed_vertices: torch.Tensor, + reference_vertices: torch.Tensor, + max_correspondence_distance: float = 0.1, ) -> Tuple[torch.Tensor, torch.Tensor]: r"""For each point in input, find nearest neighbor index in target. Args: - input (torch.Tensor): Input of shape (M, 3) or (B, M, 3). - target (torch.Tensor): Target of shape (N, 3) or (B, N, 3). - max_distance (float): Maximum distance between point correspondences. + observed_vertices (torch.Tensor): Observed vertices of shape (M, 3) or (B, M, 3). + reference_vertices (torch.Tensor): Reference vertices of shape (N, 3) or (B, N, 3). + max_correspondence_distance (float): Maximum distance between point correspondences. Returns: Tuple[torch.Tensor,torch.Tensor]: - Match-indices of shape (M) or (B, M), where mi is the index of the nearest neighbor in target. - Mask of shape (M) or (B, M). """ - if input.shape[-1] != 3 or target.shape[-1] != 3: + if observed_vertices.shape[-1] != 3 or reference_vertices.shape[-1] != 3: raise ValueError("Input and target must have shape (..., 3).") - if max_distance < 0: + if max_correspondence_distance < 0: raise ValueError("Max distance must be positive.") - distances = torch.cdist(input, target, p=2) # (M, N) - min_distance, matchindices = torch.min(distances, dim=-1) # (M) - mask = min_distance < max_distance - return matchindices, mask + distances = torch.cdist(observed_vertices, reference_vertices, p=2) # (M, N) + min_distance, match_indices = torch.min(distances, dim=-1) # (M) + mask = min_distance < max_correspondence_distance + return match_indices, mask def centroid_alignment( - Xs: List[torch.Tensor], - Ys: List[torch.Tensor], + observed_vertices: List[torch.Tensor], + reference_vertices: List[torch.Tensor], ) -> torch.Tensor: - r"""Aligns centroids of Xs and Ys as an initial guess. + r"""Aligns centroids of observed_vertices and Ys as an initial guess. Args: - Xs (List[torch.Tensor]): List of poinclouds of shape (Mi, 3). - Ys (List[torch.Tensor]): List of pointclouds of shape (Ni, 3). + observed_vertices (List[torch.Tensor]): List of poinclouds of shape (Mi, 3). + reference_vertices (List[torch.Tensor]): List of pointclouds of shape (Ni, 3). Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ Xs = Ys. + torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ observed_vertices = reference_vertices. """ # for each cloud compute centroid - Xs_centroids = [torch.mean(observation, dim=-2) for observation in Xs] - Ys_centroids = [torch.mean(mesh, dim=-2) for mesh in Ys] + observed_centroids = [ + torch.mean(observation, dim=-2) for observation in observed_vertices + ] + reference_centroids = [torch.mean(mesh, dim=-2) for mesh in reference_vertices] # estimate transform R, t = kabsch_register( - torch.stack(Xs_centroids).unsqueeze(0), - torch.stack(Ys_centroids).unsqueeze(0), + observed_vertices=torch.stack(observed_centroids).unsqueeze(0), + reference_vertices=torch.stack(reference_centroids).unsqueeze(0), ) HT = torch.eye(4, dtype=R.dtype, device=R.device) @@ -103,63 +108,70 @@ def centroid_alignment( def point_to_point_icp( HT_init: torch.Tensor, - observations: List[torch.Tensor], - meshes: List[torch.Tensor], - max_distance: float = 0.1, - max_iter: int = 100, - rmse_change: float = 1e-6, - exit_early: bool = True, -) -> torch.Tensor: + observed_vertices: List[torch.Tensor], + reference_vertices: List[torch.Tensor], + max_correspondence_distance: float = 0.1, + max_iterations: int = 100, + rmse_change_tolerance: float = 1e-6, +) -> RegistrationResult: r"""Hydra iterative closest point algorithm. Args: - HT_init: Initial guess. HT_init @ observations = meshes. - observations: List of observations of shape (Mi, 3). - meshes: List of meshes of shape (Ni, 3). - max_distance: Maximum distance between point correspondences. - max_iter: Maximum number of iterations. - rmse_change: Minimum change in rmse to continue iterating. + HT_init: Initial guess. HT_init @ observed_vertices = reference_vertices. + observed_vertices: List of observed vertices of shape (Mi, 3). + reference_vertices: List of reference vertices of shape (Ni, 3). + max_correspondence_distance: Maximum distance between point correspondences. + max_iterations: Maximum number of iterations. + rmse_change_tolerance: Minimum change in rmse to continue iterating. Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ observations = meshes. + RegistrationResult: Result with homogeneous transformation of shape (4, 4). HT @ observed_vertices = reference_vertices. """ HT = HT_init # registration - prev_rmse = float("inf") - for _ in track(range(max_iter), description=f"Running Hydra ICP..."): - observation_corr = [] - mesh_corr = [] - for i in range(len(meshes)): + previous_rmse = float("inf") + for iteration in track( + range(max_iterations), description=f"Running point to point ICP..." + ): + observed_correspondences = [] + reference_correspondences = [] + for i in range(len(reference_vertices)): # search correspondences - observations_tf = observations[i] @ HT[:3, :3].T + HT[:3, 3] - matchindices, mask = correspondence_indices( - observations_tf, meshes[i], max_distance + observations_tf = observed_vertices[i] @ HT[:3, :3].T + HT[:3, 3] + match_indices, mask = correspondence_indices( + observations_tf, reference_vertices[i], max_correspondence_distance ) - observation_corr.append(observations[i][mask]) - mesh_corr.append(meshes[i][matchindices[mask]].squeeze()) + observed_correspondences.append(observed_vertices[i][mask]) + reference_correspondences.append( + reference_vertices[i][match_indices[mask]].squeeze() + ) - observation_corr = torch.concatenate(observation_corr).unsqueeze(0) - mesh_corr = torch.concatenate(mesh_corr).unsqueeze(0) + observed_correspondences = torch.concatenate( + observed_correspondences + ).unsqueeze(0) + reference_correspondences = torch.concatenate( + reference_correspondences + ).unsqueeze(0) ( R, t, ) = kabsch_register( - observation_corr, - mesh_corr, + observed_correspondences, + reference_correspondences, ) R = R.squeeze(0) t = t.squeeze(0) HT[:3, :3] = R.T HT[:3, 3] = t - # compute rmse between observation and mesh_corr + # compute rmse between observed_correspondences and reference_correspondences rmse = torch.sqrt( torch.mean( torch.sum( torch.pow( - mesh_corr - observation_corr, + reference_correspondences - observed_correspondences, 2, ), dim=-1, @@ -167,100 +179,115 @@ def point_to_point_icp( ) ) - if abs(prev_rmse - rmse.item()) < rmse_change and exit_early: - print("Converged early. Exiting.") - break + if abs(previous_rmse - rmse.item()) < rmse_change_tolerance: + return RegistrationResult( + extrinsics=HT, + iterations=iteration, + termination_reason=TerminationReason.CONVERGED, + ) - prev_rmse = rmse.item() + previous_rmse = rmse.item() - print("HT estimate:\n", HT) - return HT + return RegistrationResult( + extrinsics=HT, + iterations=max_iterations, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) def point_to_plane_robust_icp( HT_init: torch.Tensor, - observations: List[torch.Tensor], - meshes: List[torch.Tensor], - mesh_normals: List[torch.Tensor], - max_distance: float = 0.1, - outer_max_iter: int = 100, - inner_max_iter: int = 3, - rmse_change: float = 1e-6, -) -> torch.Tensor: + observed_vertices: List[torch.Tensor], + reference_vertices: List[torch.Tensor], + reference_normals: List[torch.Tensor], + max_correspondence_distance: float = 0.1, + max_outer_iterations: int = 100, + max_inner_iterations: int = 3, + rmse_change_tolerance: float = 1e-6, +) -> RegistrationResult: r"""Lie-algebra point-to-plane ICP with robust loss, refer to section 1 https://drive.google.com/file/d/1iIUqKchAbcYzwyS2D6jNI1J6KotReD1h/view?usp=sharing. Args: - HT_init: Initial guess. HT_init @ observations = meshes. - observations: List of observations of shape (Mi, 3). - meshes: List of meshes of shape (Ni, 3). - mesh_normals: List of mesh normals of shape (Ni, 3). - max_distance: Maximum distance between point correspondences. - outer_max_iter: Maximum number of outer iterations. - inner_max_iter: Maximum number of inner iterations. - rmse_change: Minimum change in rmse to continue iterating. + HT_init: Initial guess. HT_init @ observed_vertices = reference_vertices. + observed_vertices: List of observed vertices of shape (Mi, 3). + reference_vertices: List of reference vertices of shape (Ni, 3). + reference_normals: List of reference normals of shape (Ni, 3). + max_correspondence_distance: Maximum distance between point correspondences. + max_outer_iterations: Maximum number of outer iterations. + max_inner_iterations: Maximum number of inner iterations. + rmse_change_tolerance: Minimum change in rmse to continue iterating. Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ observations = meshes. + RegistrationResult: Result with homogeneous transformation of shape (4, 4). HT @ observed_vertices = reference_vertices. """ - HT = HT_init # HT @ observation = mesh + HT = HT_init # HT @ observed_vertices = reference_vertices - observations_cross_mat = [] - for i in range(len(observations)): + observed_cross_mat = [] + for i in range(len(observed_vertices)): # build observation cross product matrix, refer eq. 4 (gets created once) - observations_cross_mat.append( + observed_cross_mat.append( torch.stack( [ - torch.zeros_like(observations[i][:, 0]), - -observations[i][:, 2], - observations[i][:, 1], - observations[i][:, 2], - torch.zeros_like(observations[i][:, 0]), - -observations[i][:, 0], - -observations[i][:, 1], - observations[i][:, 0], - torch.zeros_like(observations[i][:, 0]), + torch.zeros_like(observed_vertices[i][:, 0]), + -observed_vertices[i][:, 2], + observed_vertices[i][:, 1], + observed_vertices[i][:, 2], + torch.zeros_like(observed_vertices[i][:, 0]), + -observed_vertices[i][:, 0], + -observed_vertices[i][:, 1], + observed_vertices[i][:, 0], + torch.zeros_like(observed_vertices[i][:, 0]), ], dim=-1, ).reshape(-1, 3, 3) ) # implementation of algorithm 1 - prev_rmse = float("inf") + previous_rmse = float("inf") dTh = torch.zeros_like(HT) - for _ in track(range(outer_max_iter), description=f"Running Hydra robust ICP..."): - observations_corr = [] - observations_cross_mat_corr = [] - meshes_corr = [] - meshes_normals_corr = [] - - for i in range(len(observations)): - if len(observations) != len(meshes): + for outer_iteration in track( + range(max_outer_iterations), description=f"Running point to plane robust ICP..." + ): + observed_correspondences = [] + observed_cross_mat_correspondences = [] + reference_correspondences = [] + reference_normals_correspondences = [] + + for i in range(len(observed_vertices)): + if len(observed_vertices) != len(reference_vertices): raise ValueError("Length of observations and meshes must be the same.") # search correspondences - observations_tf = observations[i] @ HT[:3, :3].T + HT[:3, 3] - matchindices, mask = correspondence_indices( - observations_tf, meshes[i], max_distance + observed_vertices_tf = observed_vertices[i] @ HT[:3, :3].T + HT[:3, 3] + match_indices, mask = correspondence_indices( + observed_vertices_tf, reference_vertices[i], max_correspondence_distance ) - observations_corr.append(observations[i][mask]) - observations_cross_mat_corr.append(observations_cross_mat[i][mask]) - meshes_corr.append(meshes[i][matchindices[mask].squeeze()]) - meshes_normals_corr.append(mesh_normals[i][matchindices[mask].squeeze()]) + observed_correspondences.append(observed_vertices[i][mask]) + observed_cross_mat_correspondences.append(observed_cross_mat[i][mask]) + reference_correspondences.append( + reference_vertices[i][match_indices[mask].squeeze()] + ) + reference_normals_correspondences.append( + reference_normals[i][match_indices[mask].squeeze()] + ) - observations_corr = torch.cat(observations_corr) - observations_cross_mat_corr = torch.cat(observations_cross_mat_corr) - meshes_corr = torch.cat(meshes_corr) - meshes_normals_corr = torch.cat(meshes_normals_corr) + observed_correspondences = torch.cat(observed_correspondences) + observed_cross_mat_correspondences = torch.cat( + observed_cross_mat_correspondences + ) + reference_correspondences = torch.cat(reference_correspondences) + reference_normals_correspondences = torch.cat(reference_normals_correspondences) - for _ in range(inner_max_iter): + for _ in range(max_inner_iterations): # ||A @ dTh - B||^2, refer eq. 14 - Al = meshes_normals_corr @ HT[:3, :3] # eq. 18 - Au = -Al.unsqueeze(1) @ observations_cross_mat_corr # eq. 19 + Al = reference_normals_correspondences @ HT[:3, :3] # eq. 18 + Au = -Al.unsqueeze(1) @ observed_cross_mat_correspondences # eq. 19 A = torch.cat((Au.squeeze(), Al.squeeze()), dim=-1) B = torch.linalg.vecdot( - meshes_normals_corr, - meshes_corr - (observations_corr @ HT[:3, :3].T + HT[:3, 3]), + reference_normals_correspondences, + reference_correspondences + - (observed_correspondences @ HT[:3, :3].T + HT[:3, 3]), ) # weight associated with Huber loss kappa = ( @@ -286,12 +313,12 @@ def point_to_plane_robust_icp( HT = HT @ torch.linalg.matrix_exp(dTh) - # compute rmse between observation and mesh_corr + # compute rmse between observation and mesh_correspondences rmse = torch.sqrt( torch.mean( torch.sum( torch.pow( - meshes_corr - observations_corr, + reference_correspondences - observed_correspondences, 2, ), dim=-1, @@ -299,11 +326,17 @@ def point_to_plane_robust_icp( ) ) - if abs(prev_rmse - rmse.item()) < rmse_change: - print("Converged early. Exiting.") - break + if abs(previous_rmse - rmse.item()) < rmse_change_tolerance: + return RegistrationResult( + extrinsics=HT, + iterations=outer_iteration, + termination_reason=TerminationReason.CONVERGED, + ) - prev_rmse = rmse.item() + previous_rmse = rmse.item() - print("HT estimate:\n", HT) - return HT + return RegistrationResult( + extrinsics=HT, + iterations=max_outer_iterations, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) diff --git a/roboreg/registration/point_cloud/solver.py b/roboreg/registration/point_cloud/solver.py index 482d26a..0a36515 100644 --- a/roboreg/registration/point_cloud/solver.py +++ b/roboreg/registration/point_cloud/solver.py @@ -1,5 +1,5 @@ from dataclasses import dataclass -from typing import List, Optional +from typing import Callable, List, Optional import numpy as np import torch @@ -21,18 +21,24 @@ @dataclass(frozen=True) -class _HydraProblem: +class HydraProblem: observed_vertices: List[torch.Tensor] reference_vertices: List[torch.Tensor] reference_normals: Optional[List[torch.Tensor]] = None +HydraCallback = Callable[ + ["HydraProblem", RegistrationResult], + None, +] + + def _prepare_hydra_problem( request: HydraRequest, config: HydraConfig, device: torch.device, compute_normals: bool = True, -) -> _HydraProblem: +) -> HydraProblem: # 1) construct robot on request robot = Robot.from_robot_data( robot_data=request.robot_data, @@ -56,9 +62,9 @@ def _prepare_hydra_problem( xyzs = depth_to_xyz( depth=depths, intrinsics=intrinsics, - z_min=config.observation.z_min, - z_max=config.observation.z_max, - conversion_factor=config.observation.depth_conversion_factor, + z_min=config.depth_to_point_cloud.z_min, + z_max=config.depth_to_point_cloud.z_max, + conversion_factor=config.depth_to_point_cloud.depth_conversion_factor, ) height, width = request.observations.shape xyzs = xyzs.view(-1, height * width, 3) # flatten BxHxWx3 -> Bx(H*W)x3 @@ -79,18 +85,18 @@ def _prepare_hydra_problem( mask, dilation_kernel=np.ones( [ - config.observation.dilation_kernel_size, - config.observation.dilation_kernel_size, + config.depth_to_point_cloud.dilation_kernel_size, + config.depth_to_point_cloud.dilation_kernel_size, ] ), erosion_kernel=np.ones( [ - config.observation.erosion_kernel_size, - config.observation.erosion_kernel_size, + config.depth_to_point_cloud.erosion_kernel_size, + config.depth_to_point_cloud.erosion_kernel_size, ] ), ) - if config.observation.use_mask_boundary + if config.depth_to_point_cloud.use_mask_boundary else mask ), ), @@ -134,7 +140,7 @@ def _prepare_hydra_problem( if mesh_normals is not None: mesh_normals[i] = mesh_normals[i][idx] - return _HydraProblem( + return HydraProblem( observed_vertices=observed_vertices, reference_vertices=mesh_vertices, reference_normals=mesh_normals, @@ -146,9 +152,11 @@ def __init__( self, config: HydraICPConfig | None = None, device: torch.device | str = "cuda", + callback: HydraCallback | None = None, ) -> None: self._config = config or HydraICPConfig() self._device = torch.device(device) + self._callback = callback def __call__(self, request: HydraRequest) -> RegistrationResult: hydra_problem = _prepare_hydra_problem( @@ -160,17 +168,17 @@ def __call__(self, request: HydraRequest) -> RegistrationResult: HT_init = centroid_alignment( hydra_problem.observed_vertices, hydra_problem.reference_vertices ) - HT = point_to_point_icp( + result = point_to_point_icp( HT_init, hydra_problem.observed_vertices, hydra_problem.reference_vertices, - max_distance=self._config.hydra.max_correspondence_distance, - max_iter=self._config.max_iterations, - rmse_change=self._config.hydra.rmse_change_tolerance, - ) - return RegistrationResult( - extrinsics=HT, + max_correspondence_distance=self._config.hydra.max_correspondence_distance, + max_iterations=self._config.max_iterations, + rmse_change_tolerance=self._config.hydra.rmse_change_tolerance, ) + if self._callback is not None: + self._callback(hydra_problem, result) + return result class HydraRobustICP: @@ -178,9 +186,11 @@ def __init__( self, config: HydraRobustICPConfig | None = None, device: torch.device | str = "cuda", + callback: HydraCallback | None = None, ) -> None: self._config = config or HydraRobustICPConfig() self._device = torch.device(device) + self._callback = callback def __call__(self, request: HydraRequest) -> RegistrationResult: hydra_problem = _prepare_hydra_problem( @@ -192,16 +202,16 @@ def __call__(self, request: HydraRequest) -> RegistrationResult: HT_init = centroid_alignment( hydra_problem.observed_vertices, hydra_problem.reference_vertices ) - HT = point_to_plane_robust_icp( + result = point_to_plane_robust_icp( HT_init, hydra_problem.observed_vertices, hydra_problem.reference_vertices, hydra_problem.reference_normals, - max_distance=self._config.hydra.max_correspondence_distance, - outer_max_iter=self._config.outer_max_iterations, - inner_max_iter=self._config.inner_max_iterations, - rmse_change=self._config.hydra.rmse_change_tolerance, - ) - return RegistrationResult( - extrinsics=HT, + max_correspondence_distance=self._config.hydra.max_correspondence_distance, + max_outer_iterations=self._config.max_outer_iterations, + max_inner_iterations=self._config.max_inner_iterations, + rmse_change_tolerance=self._config.hydra.rmse_change_tolerance, ) + if self._callback is not None: + self._callback(hydra_problem, result) + return result diff --git a/test/test_hydra_icp.py b/test/test_hydra_icp.py index 5e1ee86..42a9b79 100644 --- a/test/test_hydra_icp.py +++ b/test/test_hydra_icp.py @@ -76,19 +76,23 @@ def test_index_shape( raise ValueError("Indices contain negative indices.") # single input - input = torch.rand(M, dim) - target = torch.rand(N, dim) # e.g. the mesh vertices + observed_vertices = torch.rand(M, dim) + reference_vertices = torch.rand(N, dim) # e.g. the mesh vertices matchindices, mask = correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 # remove some elements randomly + observed_vertices, + reference_vertices, + max_correspondence_distance=np.sqrt(dim) / 2.0, # remove some elements randomly ) test_index_shape(matchindices, mask, torch.Size([M]), N) # batched input batch_size = 2 - input = torch.rand(batch_size, M, dim) - target = torch.rand(batch_size, N, dim) + observed_vertices = torch.rand(batch_size, M, dim) + reference_vertices = torch.rand(batch_size, N, dim) matchindices, mask = correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 + observed_vertices, + reference_vertices, + max_correspondence_distance=np.sqrt(dim) / 2.0, ) test_index_shape(matchindices, mask, torch.Size([batch_size, M]), N) @@ -96,10 +100,12 @@ def test_index_shape( M = 10 N = 100 - input = torch.rand(M, dim) - target = torch.rand(N, dim) + observed_vertices = torch.rand(M, dim) + reference_vertices = torch.rand(N, dim) matchindices, mask = correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 + observed_vertices, + reference_vertices, + max_correspondence_distance=np.sqrt(dim) / 2.0, ) test_index_shape(matchindices, mask, torch.Size([M]), N) @@ -156,19 +162,19 @@ def test_hydra_point_to_point_icp(): ) # perform forward kinematics - mesh_vertices = meshes.vertices.clone() + reference_vertices = meshes.vertices.clone() joint_states = torch.tensor( np.array(observations.joint_states), dtype=torch.float32, device=device ) ht_lookup = kinematics.forward_kinematics(joint_states) for link_name, ht in ht_lookup.items(): - mesh_vertices[ + reference_vertices[ :, meshes.lower_vertex_index_lookup[ link_name ] : meshes.upper_vertex_index_lookup[link_name], ] = torch.matmul( - mesh_vertices[ + reference_vertices[ :, meshes.lower_vertex_index_lookup[ link_name @@ -176,7 +182,7 @@ def test_hydra_point_to_point_icp(): ], ht.transpose(-1, -2), ) - mesh_vertices = from_homogeneous(mesh_vertices) + reference_vertices = from_homogeneous(reference_vertices) # turn depths into xyzs intrinsics = torch.tensor(intrinsics, dtype=torch.float32, device=device) @@ -196,8 +202,8 @@ def test_hydra_point_to_point_icp(): xyzs = xyzs.view(-1, height, width, 3) xyzs = [xyz.squeeze() for xyz in xyzs.cpu().numpy()] - # mesh vertices to list - mesh_vertices = [mesh_vertices[i].contiguous() for i in range(batch_size)] + # reference vertices to list + reference_vertices = [reference_vertices[i].contiguous() for i in range(batch_size)] # clean observed vertices and turn into tensor observed_vertices = [ @@ -211,30 +217,32 @@ def test_hydra_point_to_point_icp(): # sample 5000 points per mesh for i in range(batch_size): - idx = torch.randperm(mesh_vertices[i].shape[0])[:5000] - mesh_vertices[i] = mesh_vertices[i][idx] + idx = torch.randperm(reference_vertices[i].shape[0])[:5000] + reference_vertices[i] = reference_vertices[i][idx] - HT_init = centroid_alignment(observed_vertices, mesh_vertices) - HT = point_to_point_icp( + HT_init = centroid_alignment(observed_vertices, reference_vertices) + registration_result = point_to_point_icp( HT_init, observed_vertices, - mesh_vertices, - max_distance=0.1, - max_iter=int(1e3), - rmse_change=1e-8, + reference_vertices, + max_correspondence_distance=0.1, + max_iterations=int(1e3), + rmse_change_tolerance=1e-8, ) # visualize visualizer = RegistrationVisualizer() - visualizer(mesh_vertices=mesh_vertices, observed_vertices=observed_vertices) + visualizer(mesh_vertices=reference_vertices, observed_vertices=observed_vertices) visualizer( - mesh_vertices=mesh_vertices, + mesh_vertices=reference_vertices, observed_vertices=observed_vertices, - HT=torch.linalg.inv(HT), + HT=torch.linalg.inv(registration_result.extrinsics), ) # to numpy - np.save(os.path.join(path, "HT_hydra.npy"), HT.cpu().numpy()) + np.save( + os.path.join(path, "HT_hydra.npy"), registration_result.extrinsics.cpu().numpy() + ) @pytest.mark.skip(reason="To be fixed.") @@ -289,19 +297,19 @@ def test_hydra_point_to_plane_robust_icp() -> None: ) # perform forward kinematics - mesh_vertices = meshes.vertices.clone() + reference_vertices = meshes.vertices.clone() joint_states = torch.tensor( np.array(observations.joint_states), dtype=torch.float32, device=device ) ht_lookup = kinematics.forward_kinematics(joint_states) for link_name, ht in ht_lookup.items(): - mesh_vertices[ + reference_vertices[ :, meshes.lower_vertex_index_lookup[ link_name ] : meshes.upper_vertex_index_lookup[link_name], ] = torch.matmul( - mesh_vertices[ + reference_vertices[ :, meshes.lower_vertex_index_lookup[ link_name @@ -329,12 +337,12 @@ def test_hydra_point_to_plane_robust_icp() -> None: xyzs = [xyz.squeeze() for xyz in xyzs.cpu().numpy()] # mesh vertices to list - mesh_vertices = from_homogeneous(mesh_vertices) - mesh_vertices = [mesh_vertices[i].contiguous() for i in range(batch_size)] - mesh_normals = [] + reference_vertices = from_homogeneous(reference_vertices) + reference_vertices = [reference_vertices[i].contiguous() for i in range(batch_size)] + reference_normals = [] for i in range(batch_size): - mesh_normals.append( - compute_vertex_normals(vertices=mesh_vertices[i], faces=meshes.faces) + reference_normals.append( + compute_vertex_normals(vertices=reference_vertices[i], faces=meshes.faces) ) # clean observed vertices and turn into tensor @@ -349,33 +357,35 @@ def test_hydra_point_to_plane_robust_icp() -> None: # sample 5000 points per mesh for i in range(batch_size): - idx = torch.randperm(mesh_vertices[i].shape[0])[:5000] - mesh_vertices[i] = mesh_vertices[i][idx] - mesh_normals[i] = mesh_normals[i][idx] + idx = torch.randperm(reference_vertices[i].shape[0])[:5000] + reference_vertices[i] = reference_vertices[i][idx] + reference_normals[i] = reference_normals[i][idx] - HT_init = centroid_alignment(observed_vertices, mesh_vertices) - HT = point_to_plane_robust_icp( + HT_init = centroid_alignment(observed_vertices, reference_vertices) + registration_result = point_to_plane_robust_icp( HT_init, observed_vertices, - mesh_vertices, - mesh_normals, - max_distance=0.1, - outer_max_iter=int(50), - inner_max_iter=10, + reference_vertices, + reference_normals, + max_correspondence_distance=0.1, + max_outer_iterations=50, + max_inner_iterations=10, ) # visualize visualizer = RegistrationVisualizer() - visualizer(mesh_vertices=mesh_vertices, observed_vertices=observed_vertices) + visualizer(mesh_vertices=reference_vertices, observed_vertices=observed_vertices) visualizer( - mesh_vertices=mesh_vertices, + mesh_vertices=reference_vertices, observed_vertices=observed_vertices, - HT=torch.linalg.inv(HT), + HT=torch.linalg.inv(registration_result.extrinsics), ) # to numpy - HT = HT.cpu().numpy() - np.save(os.path.join(path, "HT_hydra_robust.npy"), HT) + np.save( + os.path.join(path, "HT_hydra_robust.npy"), + registration_result.extrinsics.cpu().numpy(), + ) if __name__ == "__main__": From 7cf5850e8a2046d4f4f1b6f287cdd43865115b36 Mon Sep 17 00:00:00 2001 From: mhubii Date: Sat, 1 Aug 2026 17:03:20 +0100 Subject: [PATCH 05/12] update mask utilities (#74) --- roboreg/util/mask.py | 97 +++++++++++++++++++++++++++++++++--------- test/util/test_mask.py | 51 +++++++++------------- 2 files changed, 97 insertions(+), 51 deletions(-) diff --git a/roboreg/util/mask.py b/roboreg/util/mask.py index 63ecf3e..d90f450 100644 --- a/roboreg/util/mask.py +++ b/roboreg/util/mask.py @@ -1,48 +1,103 @@ import cv2 import numpy as np -from scipy.signal import convolve2d + + +def _as_binary_uint8(mask: np.ndarray) -> np.ndarray: + if mask.ndim != 2: + raise ValueError(f"Expected a 2D mask, got shape {mask.shape}.") + + return np.where(mask > 0, 255, 0).astype(np.uint8) + + +def _as_uint8_kernel(kernel: np.ndarray) -> np.ndarray: + if kernel.ndim != 2: + raise ValueError(f"Expected a 2D kernel, got shape {kernel.shape}.") + + return np.where(kernel > 0, 1, 0).astype(np.uint8) def mask_dilate_with_kernel( - mask: np.ndarray, kernel: np.ndarray = np.ones([10, 10]) + mask: np.ndarray, + kernel: np.ndarray | None = None, ) -> np.ndarray: - extended_mask = convolve2d(mask, kernel, mode="same") - extended_mask = np.where(extended_mask > 0.0, 255.0, 0.0).astype(np.uint8) - return extended_mask + if kernel is None: + kernel = np.ones((10, 10), dtype=np.uint8) + + mask = _as_binary_uint8(mask) + kernel = _as_uint8_kernel(kernel) + + return cv2.dilate(mask, kernel) def mask_distance_transform(mask: np.ndarray) -> np.ndarray: + mask = _as_binary_uint8(mask) + return cv2.distanceTransform(mask, cv2.DIST_L2, cv2.DIST_MASK_PRECISE) def mask_erode_with_kernel( - mask: np.ndarray, kernel: np.ndarray = np.ones([4, 4]) + mask: np.ndarray, + kernel: np.ndarray | None = None, +) -> np.ndarray: + if kernel is None: + kernel = np.ones((4, 4), dtype=np.uint8) + + mask = _as_binary_uint8(mask) + kernel = _as_uint8_kernel(kernel) + + return cv2.erode(mask, kernel) + + +def mask_exponential_decay( + mask: np.ndarray, + sigma: float = 2.0, ) -> np.ndarray: - shrinked_mask = cv2.erode(mask, kernel) - return shrinked_mask + if sigma <= 0: + raise ValueError("sigma must be positive.") + mask = _as_binary_uint8(mask) + inverse_mask = cv2.bitwise_not(mask) -def mask_exponential_decay(mask: np.ndarray, sigma: float = 2.0) -> np.ndarray: - inverse_mask = np.where(mask > 0.0, 0.0, 1.0).astype(np.uint8) distance_map = mask_distance_transform(inverse_mask) - distance_map = np.exp(-distance_map / sigma) - return distance_map + + return np.exp(-distance_map / sigma).astype(np.float32) def mask_extract_boundary( mask: np.ndarray, - erosion_kernel: np.ndarray = np.ones([10, 10]), + erosion_kernel: np.ndarray | None = None, ) -> np.ndarray: - boundary_mask = mask - cv2.erode(mask, erosion_kernel) - return boundary_mask + if erosion_kernel is None: + erosion_kernel = np.ones((10, 10), dtype=np.uint8) + + mask = _as_binary_uint8(mask) + eroded_mask = mask_erode_with_kernel( + mask=mask, + kernel=erosion_kernel, + ) + + return cv2.subtract(mask, eroded_mask) def mask_extract_extended_boundary( mask: np.ndarray, - dilation_kernel: np.ndarray = np.ones([10, 10]), - erosion_kernel: np.ndarray = np.ones([10, 10]), + dilation_kernel: np.ndarray | None = None, + erosion_kernel: np.ndarray | None = None, ) -> np.ndarray: - extended_boundary_mask = mask_dilate_with_kernel( - mask=mask, kernel=dilation_kernel - ) - mask_erode_with_kernel(mask=mask, kernel=erosion_kernel) - return extended_boundary_mask + if dilation_kernel is None: + dilation_kernel = np.ones((10, 10), dtype=np.uint8) + + if erosion_kernel is None: + erosion_kernel = np.ones((10, 10), dtype=np.uint8) + + dilated_mask = mask_dilate_with_kernel( + mask=mask, + kernel=dilation_kernel, + ) + + eroded_mask = mask_erode_with_kernel( + mask=mask, + kernel=erosion_kernel, + ) + + return cv2.subtract(dilated_mask, eroded_mask) diff --git a/test/util/test_mask.py b/test/util/test_mask.py index 651573c..3686e25 100644 --- a/test/util/test_mask.py +++ b/test/util/test_mask.py @@ -26,10 +26,8 @@ def test_dilate_with_kernel() -> None: cv2.IMREAD_GRAYSCALE, ) dilated_mask = mask_dilate_with_kernel(mask) - cv2.imshow("mask", mask) - cv2.imshow("dilated_mask", dilated_mask) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_dilate_with_kernel.mask.png", mask) + cv2.imwrite("test_dilate_with_kernel.dilated_mask.png", dilated_mask) @pytest.mark.skip(reason="To be fixed.") @@ -45,10 +43,8 @@ def test_distance_transform() -> None: distance_map = (distance_map / distance_map.max() * 255.0).astype( np.uint8 ) # normalize for visualization - cv2.imshow("mask", mask) - cv2.imshow("distance_map", distance_map) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_distance_transform.mask.png", mask) + cv2.imwrite("test_distance_transform.distance_map.png", distance_map) # show inverse distance map inverse_mask = np.where(mask > 0, 0, 255).astype(np.uint8) @@ -58,10 +54,10 @@ def test_distance_transform() -> None: ).astype( np.uint8 ) # normalize for visualization - cv2.imshow("inverse_mask", inverse_mask) - cv2.imshow("inverse_distance_map", inverse_distance_map) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_distance_transform.inverse_mask.png", inverse_mask) + cv2.imwrite( + "test_distance_transform.inverse_distance_map.png", inverse_distance_map + ) @pytest.mark.skip(reason="To be fixed.") @@ -72,10 +68,8 @@ def test_erode_with_kernel() -> None: cv2.IMREAD_GRAYSCALE, ) eroded_mask = mask_erode_with_kernel(mask) - cv2.imshow("mask", mask) - cv2.imshow("eroded_mask", eroded_mask) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_erode_with_kernel.mask.png", mask) + cv2.imwrite("test_erode_with_kernel.eroded_mask.png", eroded_mask) @pytest.mark.skip(reason="To be fixed.") @@ -86,10 +80,8 @@ def test_exponential_decay() -> None: cv2.IMREAD_GRAYSCALE, ) exponential_decay = mask_exponential_decay(mask) - cv2.imshow("mask", mask) - cv2.imshow("exponential_decay", exponential_decay) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_exponential_decay.mask.png", mask) + cv2.imwrite("test_exponential_decay.exponential_decay.png", exponential_decay) @pytest.mark.skip(reason="To be fixed.") @@ -102,11 +94,9 @@ def test_extract_boundary() -> None: ) boundary_mask = mask_extract_boundary(mask) overlay = overlay_mask(img, boundary_mask, mode="b", alpha=1.0, scale=1.0) - cv2.imshow("mask", mask) - cv2.imshow("boundary_mask", boundary_mask) - cv2.imshow("overlay", overlay) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_extract_boundary.mask.png", mask) + cv2.imwrite("test_extract_boundary.boundary_mask.png", boundary_mask) + cv2.imwrite("test_extract_boundary.overlay.png", overlay) @pytest.mark.skip(reason="To be fixed.") @@ -121,11 +111,12 @@ def test_extract_extended_boundary() -> None: mask, dilation_kernel=np.ones([2, 2]), erosion_kernel=np.ones([10, 10]) ) overlay = overlay_mask(img, extended_boundary_mask, mode="b", alpha=1.0, scale=1.0) - cv2.imshow("mask", mask) - cv2.imshow("extended_boundary_mask", extended_boundary_mask) - cv2.imshow("overlay", overlay) - cv2.waitKey(0) - cv2.destroyAllWindows() + cv2.imwrite("test_extract_extended_boundary.mask.png", mask) + cv2.imwrite( + "test_extract_extended_boundary.extended_boundary_mask.png", + extended_boundary_mask, + ) + cv2.imwrite("test_extract_extended_boundary.overlay.png", overlay) if __name__ == "__main__": From 53a4aae8656fad5330592ec5df712f1b6028d5c0 Mon Sep 17 00:00:00 2001 From: mhubii Date: Sat, 1 Aug 2026 17:21:01 +0100 Subject: [PATCH 06/12] updated typing --- roboreg/registration/point_cloud/solver.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/roboreg/registration/point_cloud/solver.py b/roboreg/registration/point_cloud/solver.py index 0a36515..25a095b 100644 --- a/roboreg/registration/point_cloud/solver.py +++ b/roboreg/registration/point_cloud/solver.py @@ -28,7 +28,7 @@ class HydraProblem: HydraCallback = Callable[ - ["HydraProblem", RegistrationResult], + [HydraProblem, RegistrationResult], None, ] From 0267bbe2ca24a4efc1e47a7c8c4b4dcf6262d411 Mon Sep 17 00:00:00 2001 From: mhubii Date: Sat, 1 Aug 2026 17:25:06 +0100 Subject: [PATCH 07/12] image-based registration template --- roboreg/core/structs.py | 22 +-- roboreg/registration/image/config.py | 46 ++++-- roboreg/registration/image/objectives.py | 170 +++++++++++++++++++++++ roboreg/registration/image/solver.py | 54 +++++++ 4 files changed, 270 insertions(+), 22 deletions(-) create mode 100644 roboreg/registration/image/objectives.py create mode 100644 roboreg/registration/image/solver.py diff --git a/roboreg/core/structs.py b/roboreg/core/structs.py index 803a375..b524e5e 100644 --- a/roboreg/core/structs.py +++ b/roboreg/core/structs.py @@ -311,22 +311,22 @@ class VirtualCamera(Camera): - https://stackoverflow.com/questions/22064084/how-to-create-perspective-projection-matrix-given-focal-points-and-camera-princ """ - __slots__ = ["_perspective_projection", "_zmin", "_zmax"] + __slots__ = ["_perspective_projection", "_z_min", "_z_max"] def __init__( self, resolution: Tuple[int, int], intrinsics: Optional[Union[torch.FloatTensor, np.ndarray]] = None, extrinsics: Optional[Union[torch.FloatTensor, np.ndarray]] = None, - zmin: float = 0.1, - zmax: float = 100.0, + z_min: float = 0.1, + z_max: float = 100.0, device: Union[torch.device, str] = "cuda", ) -> None: super().__init__(resolution, intrinsics, extrinsics, device) # build perspective projection matrix - self._zmin = zmin - self._zmax = zmax + self._z_min = z_min + self._z_max = z_max if ( self._intrinsics.ndim == 2 @@ -351,8 +351,8 @@ def __init__( self._perspective_projection[..., 1, 2] = ( 2.0 * self._intrinsics[..., 1, 2] / self.height - 1.0 ) - self._perspective_projection[..., 2, 2] = (zmax + zmin) / (zmax - zmin) - self._perspective_projection[..., 2, 3] = 2.0 * zmax * zmin / (zmin - zmax) + self._perspective_projection[..., 2, 2] = (z_max + z_min) / (z_max - z_min) + self._perspective_projection[..., 2, 3] = 2.0 * z_max * z_min / (z_min - z_max) self._perspective_projection[..., 3, 2] = 1.0 @classmethod @@ -387,9 +387,9 @@ def perspective_projection(self) -> torch.FloatTensor: return self._perspective_projection @property - def zmin(self) -> float: - return self._zmin + def z_min(self) -> float: + return self._z_min @property - def zmax(self) -> float: - return self._zmax + def z_max(self) -> float: + return self._z_max diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py index 4fbccbf..7af110e 100644 --- a/roboreg/registration/image/config.py +++ b/roboreg/registration/image/config.py @@ -1,25 +1,30 @@ import math from dataclasses import dataclass, field +from typing import Literal, Tuple @dataclass(frozen=True) -class ConvergenceConfig: - max_iterations: int = 400 - tolerance: float = 1.0e-3 - patience: int = 50 +class CameraConfig: + z_min: float = 0.1 + z_max: float = 100.0 + target_resolution: Tuple[int, int] | None = None def __post_init__(self) -> None: - if self.max_iterations <= 0: - raise ValueError("max_iterations must be positive.") - if self.tolerance < 0: - raise ValueError("tolerance must be non-negative.") - if self.patience < 0: - raise ValueError("patience must be non-negative.") + if self.z_max <= self.z_min: + raise ValueError("z_max must be greater than z_min.") + if self.z_min < 0: + raise ValueError("z_min must be greater equal zero.") + + if self.target_resolution is not None: + height, width = self.target_resolution + + if height <= 0 or width <= 0: + raise ValueError("target_resolution dimensions must be positive.") @dataclass(frozen=True) class PlateauSchedulerConfig: - mode: str = "min" + mode: Literal["min", "max"] = "min" factor: float = 0.1 patience: int = 50 threshold: float = 1.0e-4 @@ -33,8 +38,25 @@ def __post_init__(self) -> None: raise ValueError("threshold must be non-negative.") +@dataclass(frozen=True) +class ConvergenceConfig: + max_iterations: int = 400 + tolerance: float = 1.0e-3 + patience: int = 50 + + def __post_init__(self) -> None: + if self.max_iterations <= 0: + raise ValueError("max_iterations must be positive.") + if self.tolerance < 0: + raise ValueError("tolerance must be non-negative.") + if self.patience < 0: + raise ValueError("patience must be non-negative.") + + @dataclass(frozen=True) class DRRegConfig: + camera: CameraConfig = field(default_factory=CameraConfig) + optimizer: str = "AdamW" lr: float = 3.0e-2 @@ -50,6 +72,8 @@ def __post_init__(self) -> None: @dataclass(frozen=True) class CSRegConfig: + camera: CameraConfig = field(default_factory=CameraConfig) + n_cameras: int = 50 min_distance: float = 0.5 max_distance: float = 2.0 diff --git a/roboreg/registration/image/objectives.py b/roboreg/registration/image/objectives.py new file mode 100644 index 0000000..e9614bc --- /dev/null +++ b/roboreg/registration/image/objectives.py @@ -0,0 +1,170 @@ +from dataclasses import dataclass +from typing import Protocol + +import torch + +from roboreg.losses import soft_dice_loss +from roboreg.util.mask import mask_distance_transform, mask_exponential_decay + + +class RenderingObjective(Protocol): + def validate_targets(self, targets: torch.Tensor) -> None: ... + + def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: ... + + def __call__( + self, preprocessed_targets: torch.Tensor, renders: torch.Tensor + ) -> torch.Tensor: ... + + +def _ensure_binary_masks( + targets: torch.Tensor, + threshold: float | None, +) -> torch.Tensor: + if threshold is not None: + return targets >= threshold + + is_binary = torch.logical_or( + targets == 0, + targets == 1, + ) + + if not torch.all(is_binary): + raise ValueError( + "Expected binary targets. Set a threshold to convert " + "probability maps into binary masks." + ) + + return targets.bool() + + +@dataclass(frozen=True) +class DistanceMapConfig: + threshold: float | None = None + + +class DistanceMapObjective: + r"""Computes the mean squared error between the distance transform of the target mask and the rendered mask. + Supports binary masks and probability maps as targets. A threshold is required for probability maps. + """ + + def __init__( + self, + config: DistanceMapConfig | None = None, + ) -> None: + self._config = config or DistanceMapConfig() + + def __post_init__(self) -> None: + if ( + self._config.threshold is not None + and self._config.threshold <= 0 + or self._config.threshold >= 1 + ): + raise ValueError("threshold must be in [0, 1].") + + def validate_targets(self, targets: torch.Tensor) -> None: + if not torch.all((targets >= 0) & (targets <= 1)): + raise ValueError("Expected targets in range [0, 1].") + _ensure_binary_masks(targets, threshold=self._config.threshold) + + def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: + targets = _ensure_binary_masks(targets, threshold=self._config.threshold) + targets_np = targets.cpu().numpy() + distance_maps = [mask_distance_transform(mask) for mask in targets_np] + return torch.tensor( + distance_maps, dtype=torch.float32, device=targets.device + ).unsqueeze(-1) + + def __call__( + self, preprocessed_targets: torch.Tensor, renders: torch.Tensor + ) -> torch.Tensor: + return torch.mean((preprocessed_targets - renders) ** 2) + + +@dataclass(frozen=True) +class ExponentialDecayMaskConfig: + sigma: float = 2.0 + epsilon: float = 1e-6 + threshold: float | None = None + + +class ExponentialDecayMaskObjective: + r"""Computes a soft Dice loss between an exponentially decaying target mask and the rendered mask. + Supports binary masks and probability maps as targets. A threshold is required for probability maps. + """ + + def __init__( + self, + config: ExponentialDecayMaskConfig | None = None, + ) -> None: + self._config = config or ExponentialDecayMaskConfig() + + def __post_init__(self) -> None: + if self._config.sigma <= 0: + raise ValueError("sigma must be positive.") + if self._config.epsilon <= 0: + raise ValueError("epsilon must be positive.") + if ( + self._config.threshold is not None + and self._config.threshold <= 0 + or self._config.threshold >= 1 + ): + raise ValueError("threshold must be in [0, 1].") + + def validate_targets(self, targets: torch.Tensor) -> None: + if not torch.all((targets >= 0) & (targets <= 1)): + raise ValueError("Expected targets in range [0, 1].") + _ensure_binary_masks(targets, threshold=self._config.threshold) + + def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: + targets = _ensure_binary_masks(targets, threshold=self._config.threshold) + targets_np = targets.cpu().numpy() + decay_maps = [ + mask_exponential_decay(mask, sigma=self._config.sigma) + for mask in targets_np + ] + return torch.tensor( + decay_maps, dtype=torch.float32, device=targets.device + ).unsqueeze(-1) + + def __call__( + self, preprocessed_targets: torch.Tensor, renders: torch.Tensor + ) -> torch.Tensor: + return soft_dice_loss( + preprocessed_targets, renders, epsilon=self._config.epsilon + ).mean() + + +@dataclass(frozen=True) +class ProbabilityMapConfig: + epsilon: float = 1e-6 + + +class ProbabilityMapObjective: + r"""Computes a soft Dice loss between the target probability map and the rendered mask.""" + + def __init__( + self, + config: ProbabilityMapConfig | None = None, + ) -> None: + self._config = config or ProbabilityMapConfig() + + def __post_init__(self) -> None: + if self._config.epsilon <= 0: + raise ValueError("epsilon must be positive.") + + def validate_targets(self, targets: torch.Tensor) -> None: + if not targets.is_floating_point(): + raise ValueError("Expected floating point probability targets.") + if not torch.all((targets >= 0) & (targets <= 1)): + raise ValueError("Expected targets in range [0, 1].") + + def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: + return targets + + def __call__( + self, preprocessed_targets: torch.Tensor, renders: torch.Tensor + ) -> torch.Tensor: + return soft_dice_loss( + preprocessed_targets, renders, epsilon=self._config.epsilon + ).mean() diff --git a/roboreg/registration/image/solver.py b/roboreg/registration/image/solver.py new file mode 100644 index 0000000..0c04e39 --- /dev/null +++ b/roboreg/registration/image/solver.py @@ -0,0 +1,54 @@ +import torch + +from roboreg.registration.image.config import CSRegConfig, DRRegConfig +from roboreg.registration.image.objectives import RenderingObjective +from roboreg.registration.image.request import MonocularRequest, StereoRequest +from roboreg.registration.result import RegistrationResult + + +class MonocularDiffRendRegistration: + def __init__( + self, + config: DRRegConfig, + objective: RenderingObjective, + device: torch.device | str = "cuda", + ) -> None: + self._config = config + self._objective = objective + self._device = torch.device(device) + + def __call__( + self, + request: MonocularRequest, + ) -> RegistrationResult: + pass + + +class StereoDiffRendRegistration: + def __init__( + self, + config: DRRegConfig, + objective: RenderingObjective, + device: torch.device | str = "cuda", + ) -> None: + self._config = config + self._objective = objective + self._device = torch.device(device) + + def __call__(self, request: StereoRequest) -> RegistrationResult: + pass + + +class CameraSwarmRegistration: + def __init__( + self, + config: CSRegConfig, + objective: RenderingObjective, + device: torch.device | str = "cuda", + ) -> None: + self._config = config + self._objective = objective + self._device = torch.device(device) + + def __call__(self, request: MonocularRequest) -> RegistrationResult: + pass From 0a2edd704d9091cfc49cc725675c70028052aafd Mon Sep 17 00:00:00 2001 From: mhubii Date: Sun, 2 Aug 2026 11:45:11 +0100 Subject: [PATCH 08/12] improve robustness --- roboreg/registration/image/objectives.py | 55 +++-- test/util/test_mask.py | 247 ++++++++++++++--------- 2 files changed, 178 insertions(+), 124 deletions(-) diff --git a/roboreg/registration/image/objectives.py b/roboreg/registration/image/objectives.py index e9614bc..0236829 100644 --- a/roboreg/registration/image/objectives.py +++ b/roboreg/registration/image/objectives.py @@ -1,6 +1,7 @@ from dataclasses import dataclass from typing import Protocol +import numpy as np import torch from roboreg.losses import soft_dice_loss @@ -42,6 +43,10 @@ def _ensure_binary_masks( class DistanceMapConfig: threshold: float | None = None + def __post_init__(self) -> None: + if self.threshold is not None and not 0.0 < self.threshold < 1.0: + raise ValueError("threshold must be in (0, 1).") + class DistanceMapObjective: r"""Computes the mean squared error between the distance transform of the target mask and the rendered mask. @@ -54,14 +59,6 @@ def __init__( ) -> None: self._config = config or DistanceMapConfig() - def __post_init__(self) -> None: - if ( - self._config.threshold is not None - and self._config.threshold <= 0 - or self._config.threshold >= 1 - ): - raise ValueError("threshold must be in [0, 1].") - def validate_targets(self, targets: torch.Tensor) -> None: if not torch.all((targets >= 0) & (targets <= 1)): raise ValueError("Expected targets in range [0, 1].") @@ -69,10 +66,10 @@ def validate_targets(self, targets: torch.Tensor) -> None: def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: targets = _ensure_binary_masks(targets, threshold=self._config.threshold) - targets_np = targets.cpu().numpy() + targets_np = targets.detach().cpu().numpy() distance_maps = [mask_distance_transform(mask) for mask in targets_np] - return torch.tensor( - distance_maps, dtype=torch.float32, device=targets.device + return torch.as_tensor( + np.stack(distance_maps), dtype=torch.float32, device=targets.device ).unsqueeze(-1) def __call__( @@ -87,6 +84,14 @@ class ExponentialDecayMaskConfig: epsilon: float = 1e-6 threshold: float | None = None + def __post_init__(self) -> None: + if self.sigma <= 0: + raise ValueError("sigma must be positive.") + if self.epsilon <= 0: + raise ValueError("epsilon must be positive.") + if self.threshold is not None and not 0.0 < self.threshold < 1.0: + raise ValueError("threshold must be in (0, 1).") + class ExponentialDecayMaskObjective: r"""Computes a soft Dice loss between an exponentially decaying target mask and the rendered mask. @@ -99,18 +104,6 @@ def __init__( ) -> None: self._config = config or ExponentialDecayMaskConfig() - def __post_init__(self) -> None: - if self._config.sigma <= 0: - raise ValueError("sigma must be positive.") - if self._config.epsilon <= 0: - raise ValueError("epsilon must be positive.") - if ( - self._config.threshold is not None - and self._config.threshold <= 0 - or self._config.threshold >= 1 - ): - raise ValueError("threshold must be in [0, 1].") - def validate_targets(self, targets: torch.Tensor) -> None: if not torch.all((targets >= 0) & (targets <= 1)): raise ValueError("Expected targets in range [0, 1].") @@ -118,13 +111,13 @@ def validate_targets(self, targets: torch.Tensor) -> None: def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: targets = _ensure_binary_masks(targets, threshold=self._config.threshold) - targets_np = targets.cpu().numpy() + targets_np = targets.detach().cpu().numpy() decay_maps = [ mask_exponential_decay(mask, sigma=self._config.sigma) for mask in targets_np ] - return torch.tensor( - decay_maps, dtype=torch.float32, device=targets.device + return torch.as_tensor( + np.stack(decay_maps), dtype=torch.float32, device=targets.device ).unsqueeze(-1) def __call__( @@ -139,6 +132,10 @@ def __call__( class ProbabilityMapConfig: epsilon: float = 1e-6 + def __post_init__(self) -> None: + if self.epsilon <= 0: + raise ValueError("epsilon must be positive.") + class ProbabilityMapObjective: r"""Computes a soft Dice loss between the target probability map and the rendered mask.""" @@ -149,10 +146,6 @@ def __init__( ) -> None: self._config = config or ProbabilityMapConfig() - def __post_init__(self) -> None: - if self._config.epsilon <= 0: - raise ValueError("epsilon must be positive.") - def validate_targets(self, targets: torch.Tensor) -> None: if not targets.is_floating_point(): raise ValueError("Expected floating point probability targets.") @@ -160,7 +153,7 @@ def validate_targets(self, targets: torch.Tensor) -> None: raise ValueError("Expected targets in range [0, 1].") def preprocess_targets(self, targets: torch.Tensor) -> torch.Tensor: - return targets + return targets.unsqueeze(-1) def __call__( self, preprocessed_targets: torch.Tensor, renders: torch.Tensor diff --git a/test/util/test_mask.py b/test/util/test_mask.py index 3686e25..b08a14d 100644 --- a/test/util/test_mask.py +++ b/test/util/test_mask.py @@ -1,8 +1,3 @@ -import os -import sys - -sys.path.append(os.path.join(os.path.dirname(__file__), "../..")) - import cv2 import numpy as np import pytest @@ -14,115 +9,181 @@ mask_exponential_decay, mask_extract_boundary, mask_extract_extended_boundary, - overlay_mask, ) -@pytest.mark.skip(reason="To be fixed.") -def test_dilate_with_kernel() -> None: - idx = 1 - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, +def _square_mask( + *, + size: int = 7, + start: int = 2, + end: int = 5, + dtype: np.dtype = np.uint8, + foreground_value: int | bool = 255, +) -> np.ndarray: + mask = np.zeros((size, size), dtype=dtype) + mask[start:end, start:end] = foreground_value + return mask + + +@pytest.mark.parametrize( + ("dtype", "foreground_value"), + [ + (np.bool_, True), + (np.uint8, 1), + (np.uint8, 255), + ], +) +def test_dilate_with_kernel( + dtype: np.dtype, + foreground_value: int | bool, +) -> None: + mask = _square_mask( + dtype=dtype, + foreground_value=foreground_value, ) - dilated_mask = mask_dilate_with_kernel(mask) - cv2.imwrite("test_dilate_with_kernel.mask.png", mask) - cv2.imwrite("test_dilate_with_kernel.dilated_mask.png", dilated_mask) + kernel = np.ones((3, 3), dtype=np.uint8) + result = mask_dilate_with_kernel(mask, kernel) -@pytest.mark.skip(reason="To be fixed.") -def test_distance_transform() -> None: - idx = 1 - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, - ) + expected = np.zeros((7, 7), dtype=np.uint8) + expected[1:6, 1:6] = 255 - # show distance map - distance_map = mask_distance_transform(mask) - distance_map = (distance_map / distance_map.max() * 255.0).astype( - np.uint8 - ) # normalize for visualization - cv2.imwrite("test_distance_transform.mask.png", mask) - cv2.imwrite("test_distance_transform.distance_map.png", distance_map) - - # show inverse distance map - inverse_mask = np.where(mask > 0, 0, 255).astype(np.uint8) - inverse_distance_map = mask_distance_transform(inverse_mask) - inverse_distance_map = ( - inverse_distance_map / inverse_distance_map.max() * 255.0 - ).astype( - np.uint8 - ) # normalize for visualization - cv2.imwrite("test_distance_transform.inverse_mask.png", inverse_mask) - cv2.imwrite( - "test_distance_transform.inverse_distance_map.png", inverse_distance_map - ) + np.testing.assert_array_equal(result, expected) + assert result.dtype == np.uint8 -@pytest.mark.skip(reason="To be fixed.") -def test_erode_with_kernel() -> None: - idx = 1 - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, +@pytest.mark.parametrize( + ("dtype", "foreground_value"), + [ + (np.bool_, True), + (np.uint8, 1), + (np.uint8, 255), + ], +) +def test_erode_with_kernel( + dtype: np.dtype, + foreground_value: int | bool, +) -> None: + mask = _square_mask( + start=1, + end=6, + dtype=dtype, + foreground_value=foreground_value, ) - eroded_mask = mask_erode_with_kernel(mask) - cv2.imwrite("test_erode_with_kernel.mask.png", mask) - cv2.imwrite("test_erode_with_kernel.eroded_mask.png", eroded_mask) + kernel = np.ones((3, 3), dtype=np.uint8) + + result = mask_erode_with_kernel(mask, kernel) + + expected = np.zeros((7, 7), dtype=np.uint8) + expected[2:5, 2:5] = 255 + + np.testing.assert_array_equal(result, expected) + assert result.dtype == np.uint8 + + +def test_distance_transform_single_foreground_pixel() -> None: + mask = np.zeros((5, 5), dtype=np.uint8) + mask[2, 2] = 255 + + result = mask_distance_transform(mask) + + expected = np.zeros((5, 5), dtype=np.float32) + expected[2, 2] = 1.0 + + np.testing.assert_allclose(result, expected, atol=1e-6) + assert result.dtype == np.float32 + + +def test_distance_transform_accepts_bool() -> None: + mask = np.zeros((5, 5), dtype=bool) + mask[2, 2] = True + + bool_result = mask_distance_transform(mask) + uint8_result = mask_distance_transform(mask.astype(np.uint8)) + + np.testing.assert_allclose(bool_result, uint8_result) -@pytest.mark.skip(reason="To be fixed.") def test_exponential_decay() -> None: - idx = 1 - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, - ) - exponential_decay = mask_exponential_decay(mask) - cv2.imwrite("test_exponential_decay.mask.png", mask) - cv2.imwrite("test_exponential_decay.exponential_decay.png", exponential_decay) + mask = np.zeros((7, 7), dtype=np.uint8) + mask[3, 3] = 255 + + sigma = 2.0 + result = mask_exponential_decay(mask, sigma=sigma) + + assert result.shape == mask.shape + assert result.dtype == np.float32 + assert np.all(np.isfinite(result)) + assert np.all((result >= 0.0) & (result <= 1.0)) + + # Inside the original mask, inverse-mask distance is zero. + assert result[3, 3] == pytest.approx(1.0) + + # The response should decay with distance from the mask. + assert result[3, 2] > result[3, 1] + assert result[3, 1] > result[3, 0] + + +def test_exponential_decay_rejects_invalid_sigma() -> None: + mask = np.zeros((5, 5), dtype=np.uint8) + + with pytest.raises(ValueError, match="sigma must be positive"): + mask_exponential_decay(mask, sigma=0.0) -@pytest.mark.skip(reason="To be fixed.") def test_extract_boundary() -> None: - idx = 1 - img = cv2.imread(f"test/assets/lbr_med7_r800/samples/left_image_{idx}.png") - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, + mask = _square_mask( + size=7, + start=1, + end=6, + ) + kernel = np.ones((3, 3), dtype=np.uint8) + + result = mask_extract_boundary( + mask, + erosion_kernel=kernel, ) - boundary_mask = mask_extract_boundary(mask) - overlay = overlay_mask(img, boundary_mask, mode="b", alpha=1.0, scale=1.0) - cv2.imwrite("test_extract_boundary.mask.png", mask) - cv2.imwrite("test_extract_boundary.boundary_mask.png", boundary_mask) - cv2.imwrite("test_extract_boundary.overlay.png", overlay) + + expected = np.zeros((7, 7), dtype=np.uint8) + expected[1:6, 1:6] = 255 + expected[2:5, 2:5] = 0 + + np.testing.assert_array_equal(result, expected) -@pytest.mark.skip(reason="To be fixed.") def test_extract_extended_boundary() -> None: - idx = 1 - img = cv2.imread(f"test/assets/lbr_med7_r800/samples/left_image_{idx}.png") - mask = cv2.imread( - f"test/assets/lbr_med7_r800/samples/mask_sam2_left_image_{idx}.png", - cv2.IMREAD_GRAYSCALE, + mask = _square_mask( + size=9, + start=3, + end=6, ) - extended_boundary_mask = mask_extract_extended_boundary( - mask, dilation_kernel=np.ones([2, 2]), erosion_kernel=np.ones([10, 10]) - ) - overlay = overlay_mask(img, extended_boundary_mask, mode="b", alpha=1.0, scale=1.0) - cv2.imwrite("test_extract_extended_boundary.mask.png", mask) - cv2.imwrite( - "test_extract_extended_boundary.extended_boundary_mask.png", - extended_boundary_mask, + kernel = np.ones((3, 3), dtype=np.uint8) + + result = mask_extract_extended_boundary( + mask, + dilation_kernel=kernel, + erosion_kernel=kernel, ) - cv2.imwrite("test_extract_extended_boundary.overlay.png", overlay) + dilated = cv2.dilate(mask, kernel) + eroded = cv2.erode(mask, kernel) + expected = cv2.subtract(dilated, eroded) + + np.testing.assert_array_equal(result, expected) + + +@pytest.mark.parametrize( + "function", + [ + mask_dilate_with_kernel, + mask_distance_transform, + mask_erode_with_kernel, + mask_extract_boundary, + mask_extract_extended_boundary, + ], +) +def test_mask_functions_reject_non_2d_masks(function) -> None: + mask = np.zeros((4, 4, 3), dtype=np.uint8) -if __name__ == "__main__": - test_dilate_with_kernel() - test_distance_transform() - test_erode_with_kernel() - test_exponential_decay() - test_extract_boundary() - test_extract_extended_boundary() + with pytest.raises(ValueError, match="Expected a 2D mask"): + function(mask) From 868f04b811d53e9471f1231644c272567a1f1031 Mon Sep 17 00:00:00 2001 From: mhubii Date: Sun, 2 Aug 2026 14:07:29 +0100 Subject: [PATCH 09/12] initial unified mono / stereo handling --- cli/rr_cam_swarm.py | 24 ++- cli/rr_mono_dr.py | 64 ++++---- cli/rr_stereo_dr.py | 126 +++++---------- roboreg/io/parsers.py | 221 ++++++++++++++++++-------- roboreg/registration/image/config.py | 4 +- roboreg/registration/image/request.py | 106 ++++++------ roboreg/registration/image/solver.py | 73 ++++++--- roboreg/util/transform.py | 20 ++- test/io/test_parsers.py | 74 +++++---- test/util/test_transform.py | 49 ++++++ 10 files changed, 443 insertions(+), 318 deletions(-) diff --git a/cli/rr_cam_swarm.py b/cli/rr_cam_swarm.py index ae811e0..f8d373d 100644 --- a/cli/rr_cam_swarm.py +++ b/cli/rr_cam_swarm.py @@ -6,12 +6,7 @@ import numpy as np import torch -from roboreg.core import ( - NVDiffRastRenderer, - Robot, - RobotScene, - VirtualCamera, -) +from roboreg.core import NVDiffRastRenderer, Robot, RobotScene, VirtualCamera from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, @@ -279,11 +274,15 @@ def main() -> None: ) # pre-process data + camera_name = "camera" joint_states = torch.tensor( np.array(observations.joint_states), dtype=torch.float32, device=device ) n_joint_states = joint_states.shape[0] - masks = [mask_exponential_decay(mask) for mask in observations.targets] + masks = [ + mask_exponential_decay(mask) + for mask in observations.cameras[camera_name].targets + ] masks = torch.tensor(np.array(masks), dtype=torch.float32, device=device) # scale image data (memory reduction) @@ -317,7 +316,6 @@ def main() -> None: batch_size = ( n_joint_states * args.n_cameras ) # (each camera observes n_joint_states joint states) - camera_name = "camera" camera = VirtualCamera( resolution=(height, width), intrinsics=intrinsics, @@ -366,10 +364,10 @@ def fitness_closure() -> torch.Tensor: center = particle_swarm_optimizer.particle_swarm.particles[:, 3:6] angle = particle_swarm_optimizer.particle_swarm.particles[:, -1:] extrinsics = look_at_from_angle(eye=eye, center=center, angle=angle) - scene.cameras["camera"].extrinsics = extrinsics.repeat_interleave( + scene.cameras[camera_name].extrinsics = extrinsics.repeat_interleave( n_joint_states, 0 ) - renders = scene.observe_from("camera").squeeze() + renders = scene.observe_from(camera_name).squeeze() fitness = ( soft_dice_loss(renders.unsqueeze(-1), masks.unsqueeze(-1)) .view(args.n_cameras, n_joint_states) @@ -387,12 +385,12 @@ def fitness_closure() -> torch.Tensor: current_best_render = cv2.resize( current_best_render, ( - observations.images[offset].shape[1], - observations.images[offset].shape[0], + observations.cameras[camera_name].images[offset].shape[1], + observations.cameras[camera_name].images[offset].shape[0], ), ) overlay = overlay_mask( - observations.images[offset], + observations.cameras[camera_name].images[offset], current_best_render, scale=1.0, ) diff --git a/cli/rr_mono_dr.py b/cli/rr_mono_dr.py index d9b1882..40064a8 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -17,8 +17,12 @@ load_robot_data_from_urdf_file, parse_monocular_observations, ) -from roboreg.losses import soft_dice_loss -from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask +from roboreg.registration.image.objectives import ( + DistanceMapObjective, + ExponentialDecayMaskObjective, + RenderingObjective, +) +from roboreg.util import overlay_mask from .util.validate import validate_urdf_source @@ -177,17 +181,22 @@ def main() -> None: # pre-process data joint_states = torch.tensor( - np.array(observations.joint_states), dtype=torch.float32, device=device + np.stack(observations.joint_states, axis=0), dtype=torch.float32, device=device ) + objective: RenderingObjective if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - targets = [mask_distance_transform(mask) for mask in observations.targets] + objective = DistanceMapObjective() elif mode == REGISTRATION_MODE.SEGMENTATION: - targets = [mask_exponential_decay(mask) for mask in observations.targets] + objective = ExponentialDecayMaskObjective() else: raise ValueError("Invalid registration mode.") - targets = torch.tensor( - np.array(targets), dtype=torch.float32, device=device - ).unsqueeze(-1) + preprocessed_targets = objective.preprocess_targets( + targets=torch.tensor( + np.stack(observations.cameras["camera"].targets).astype(np.float32) / 255.0, + dtype=torch.float32, + device=device, + ) + ) # instantiate camera with default identity extrinsics because we optimize for robot pose instead camera = { @@ -255,13 +264,9 @@ def main() -> None: renders = { "camera": scene.observe_from("camera"), } - if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - loss = torch.nn.functional.mse_loss(targets, renders["camera"]) - elif mode == REGISTRATION_MODE.SEGMENTATION: - loss = soft_dice_loss(targets, renders["camera"]).mean() - else: - raise ValueError("Invalid registration mode.") - + loss = objective( + preprocessed_targets=preprocessed_targets, renders=renders["camera"] + ) optimizer.zero_grad() loss.backward() optimizer.step() @@ -279,7 +284,7 @@ def main() -> None: # display optimization progress if args.display_progress: render = renders["camera"][0].squeeze().detach().cpu().numpy() - image = observations.images[0] + image = observations.cameras["camera"].images[0] render_overlay = overlay_mask( image, (render * 255.0).astype(np.uint8), @@ -288,7 +293,11 @@ def main() -> None: # difference left / right render / mask difference = ( cv2.cvtColor( - np.abs(render - observations.targets[0].astype(np.float32) / 255.0), + np.abs( + render + - observations.cameras["camera"].targets[0].astype(np.float32) + / 255.0 + ), cv2.COLOR_GRAY2BGR, ) * 255.0 @@ -296,7 +305,7 @@ def main() -> None: # overlay segmentation mask segmentation_overlay = overlay_mask( image, - observations.targets[0], + observations.cameras["camera"].targets[0], mode="b", scale=1.0, ) @@ -317,24 +326,7 @@ def main() -> None: ) cv2.waitKey(1) - # render final results and save extrinsics - with torch.no_grad(): - scene.robot.configure(joint_states, best_extrinsics_inv) - renders = scene.observe_from("camera") - - for i, render in enumerate(renders): - render = render.squeeze().cpu().numpy() - overlay = overlay_mask( - observations.images[i], (render * 255.0).astype(np.uint8), scale=1.0 - ) - difference = np.abs(render - observations.targets[i].astype(np.float32) / 255.0) - - cv2.imwrite(os.path.join(args.path, f"dr_overlay_{i}.png"), overlay) - cv2.imwrite( - os.path.join(args.path, f"dr_difference_{i}.png"), - (difference * 255.0).astype(np.uint8), - ) - + # save extrinsics np.save( os.path.join(args.path, args.output_file), best_extrinsics.cpu().numpy(), diff --git a/cli/rr_stereo_dr.py b/cli/rr_stereo_dr.py index 30de37f..b0bf448 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -10,20 +10,19 @@ import rich.progress import torch -from roboreg.core import ( - NVDiffRastRenderer, - Robot, - RobotScene, - VirtualCamera, -) +from roboreg.core import NVDiffRastRenderer, Robot, RobotScene, VirtualCamera from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, parse_stereo_observations, ) -from roboreg.losses import soft_dice_loss -from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask +from roboreg.registration.image.objectives import ( + DistanceMapObjective, + ExponentialDecayMaskObjective, + RenderingObjective, +) +from roboreg.util import overlay_mask from .util.validate import validate_urdf_source @@ -216,30 +215,29 @@ def main() -> None: # pre-process data joint_states = torch.tensor( - np.array(observations.joint_states), dtype=torch.float32, device=device + np.stack(observations.joint_states, axis=0), dtype=torch.float32, device=device ) + objective: RenderingObjective if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - left_targets = [ - mask_distance_transform(mask) for mask in observations.left_targets - ] - right_targets = [ - mask_distance_transform(mask) for mask in observations.right_targets - ] + objective = DistanceMapObjective() elif mode == REGISTRATION_MODE.SEGMENTATION: - left_targets = [ - mask_exponential_decay(mask) for mask in observations.left_targets - ] - right_targets = [ - mask_exponential_decay(mask) for mask in observations.right_targets - ] + objective = ExponentialDecayMaskObjective() else: raise ValueError("Invalid registration mode.") - left_targets = torch.tensor( - np.array(left_targets), dtype=torch.float32, device=device - ).unsqueeze(-1) - right_targets = torch.tensor( - np.array(right_targets), dtype=torch.float32, device=device - ).unsqueeze(-1) + left_preprocessed_targets = objective.preprocess_targets( + targets=torch.tensor( + np.stack(observations.cameras["left"].targets).astype(np.float32) / 255.0, + dtype=torch.float32, + device=device, + ) + ) + right_preprocessed_targets = objective.preprocess_targets( + targets=torch.tensor( + np.stack(observations.cameras["right"].targets).astype(np.float32) / 255.0, + dtype=torch.float32, + device=device, + ) + ) # instantiate: # - left camera with default identity extrinsics because we optimize for robot pose instead @@ -315,17 +313,11 @@ def main() -> None: "left": scene.observe_from("left"), "right": scene.observe_from("right"), } - if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - loss = torch.nn.functional.mse_loss( - left_targets, renders["left"] - ) + torch.nn.functional.mse_loss(right_targets, renders["right"]) - elif mode == REGISTRATION_MODE.SEGMENTATION: - loss = ( - soft_dice_loss(left_targets, renders["left"]).mean() - + soft_dice_loss(right_targets, renders["right"]).mean() - ) - else: - raise ValueError("Invalid registration mode.") + loss = objective( + preprocessed_targets=left_preprocessed_targets, renders=renders["left"] + ) + objective( + preprocessed_targets=right_preprocessed_targets, renders=renders["right"] + ) optimizer.zero_grad() loss.backward() optimizer.step() @@ -344,7 +336,7 @@ def main() -> None: if args.display_progress: render_overlays = [] left_render = renders["left"][0].squeeze().detach().cpu().numpy() - left_image = observations.left_images[0] + left_image = observations.cameras["left"].images[0] render_overlays.append( overlay_mask( left_image, @@ -353,7 +345,7 @@ def main() -> None: ) ) right_render = renders["right"][0].squeeze().detach().cpu().numpy() - right_image = observations.right_images[0] + right_image = observations.cameras["right"].images[0] render_overlays.append( overlay_mask( right_image, @@ -368,7 +360,8 @@ def main() -> None: cv2.cvtColor( np.abs( left_render - - observations.left_targets[0].astype(np.float32) / 255.0 + - observations.cameras["left"].targets[0].astype(np.float32) + / 255.0 ), cv2.COLOR_GRAY2BGR, ) @@ -380,7 +373,10 @@ def main() -> None: cv2.cvtColor( np.abs( right_render - - observations.right_targets[0].astype(np.float32) / 255.0 + - observations.cameras["right"] + .targets[0] + .astype(np.float32) + / 255.0 ), cv2.COLOR_GRAY2BGR, ) @@ -392,7 +388,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( left_image, - observations.left_targets[0], + observations.cameras["left"].targets[0], mode="b", scale=1.0, ) @@ -400,7 +396,7 @@ def main() -> None: segmentation_overlays.append( overlay_mask( right_image, - observations.right_targets[0], + observations.cameras["right"].targets[0], mode="b", scale=1.0, ) @@ -422,47 +418,7 @@ def main() -> None: ) cv2.waitKey(1) - # render final results and save extrinsics - with torch.no_grad(): - scene.robot.configure(joint_states, best_left_extrinsics_inv) - renders = { - "left": scene.observe_from("left"), - "right": scene.observe_from("right"), - } - - for i, (left_render, right_render) in enumerate( - zip(renders["left"], renders["right"]) - ): - left_render = left_render.squeeze().cpu().numpy() - right_render = right_render.squeeze().cpu().numpy() - left_overlay = overlay_mask( - observations.left_images[i], - (left_render * 255.0).astype(np.uint8), - scale=1.0, - ) - right_overlay = overlay_mask( - observations.right_images[i], - (right_render * 255.0).astype(np.uint8), - scale=1.0, - ) - left_difference = np.abs( - left_render - observations.left_targets[i].astype(np.float32) / 255.0 - ) - right_difference = np.abs( - right_render - observations.right_targets[i].astype(np.float32) / 255.0 - ) - - cv2.imwrite(os.path.join(args.path, f"left_dr_overlay_{i}.png"), left_overlay) - cv2.imwrite(os.path.join(args.path, f"right_dr_overlay_{i}.png"), right_overlay) - cv2.imwrite( - os.path.join(args.path, f"left_dr_difference_{i}.png"), - (left_difference * 255.0).astype(np.uint8), - ) - cv2.imwrite( - os.path.join(args.path, f"right_dr_difference_{i}.png"), - (right_difference * 255.0).astype(np.uint8), - ) - + # save extrinsics np.save( os.path.join(args.path, args.left_output_file), best_left_extrinsics.cpu().numpy(), diff --git a/roboreg/io/parsers.py b/roboreg/io/parsers.py index 5fe1e12..b68651f 100644 --- a/roboreg/io/parsers.py +++ b/roboreg/io/parsers.py @@ -7,7 +7,7 @@ import yaml from pytorch_kinematics import urdf_parser_py -from roboreg.registration.image.request import MonocularObservations, StereoObservations +from roboreg.registration.image.request import CameraObservations, ImageObservations from roboreg.registration.point_cloud.request import HydraObservations @@ -355,88 +355,169 @@ def parse_hydra_observations( return HydraObservations(joint_states=joint_states, masks=masks, depths=depths) -def parse_monocular_observations( - image_files: List[Path], - joint_states_files: List[Path], - target_files: List[Path], -) -> MonocularObservations: - r"""Parse monocular observations. +def _read_image(path: Path) -> np.ndarray: + image = cv2.imread(str(path), cv2.IMREAD_COLOR) - Args: - image_files (List[Path]): Image files. - joint_states_files (List[Path]): Joint states files. - target_files (List[Path]): Target files. + if image is None: + raise ValueError(f"Failed to read image '{path}'.") + + return image + + +def _read_target(path: Path) -> np.ndarray: + target = cv2.imread(str(path), cv2.IMREAD_GRAYSCALE) + + if target is None: + raise ValueError(f"Failed to read target '{path}'.") + + return target - Returns: - MonocularObservations: Data for monocular registration. - """ - if len(image_files) != len(joint_states_files) or len(image_files) != len( - target_files - ): - raise ValueError("Number of images, joint states, masks do not match.") + +def _validate_image_target_shapes( + images: list[np.ndarray], + targets: list[np.ndarray], + camera_name: str, +) -> None: + for index, (image, target) in enumerate(zip(images, targets)): + if image.shape[:2] != target.shape[:2]: + raise ValueError( + f"Camera '{camera_name}' image and target at index {index} " + f"have incompatible shapes: {image.shape[:2]} and " + f"{target.shape[:2]}." + ) + + +def parse_monocular_observations( + image_files: list[Path] | None, + joint_states_files: list[Path], + target_files: list[Path], + camera_name: str = "camera", +) -> ImageObservations: + r"""Parse monocular image-registration observations.""" + + lengths = { + "joint_states": len(joint_states_files), + "targets": len(target_files), + } + + if image_files is not None: + lengths["images"] = len(image_files) + + if len(set(lengths.values())) != 1: + raise ValueError( + f"All observation file lists must have the same length, got {lengths}." + ) + + if not joint_states_files: + raise ValueError("Expected at least one observation.") rich.print("Parsing the following files:") - rich.print(f"Images: {[f.name for f in image_files]}") - rich.print(f"Joint states: {[f.name for f in joint_states_files]}") - rich.print(f"Targets: {[f.name for f in target_files]}") + if image_files is not None: + rich.print(f"Images: {[path.name for path in image_files]}") + rich.print(f"Joint states: {[path.name for path in joint_states_files]}") + rich.print(f"Targets: {[path.name for path in target_files]}") - images = [cv2.imread(f) for f in image_files] - joint_states = [np.load(f) for f in joint_states_files] - masks = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in target_files] - if not all( - [mask.shape[:2] == image.shape[:2] for mask, image in zip(masks, images)] - ): - raise ValueError("Mask and image shapes do not match.") - return MonocularObservations( - images=images, joint_states=joint_states, targets=masks + images = ( + [_read_image(path) for path in image_files] if image_files is not None else None ) + joint_states = [np.load(path) for path in joint_states_files] + targets = [_read_target(path) for path in target_files] + + if images is not None: + _validate_image_target_shapes( + images=images, + targets=targets, + camera_name=camera_name, + ) + return ImageObservations( + joint_states=joint_states, + cameras={ + camera_name: CameraObservations( + images=images, + targets=targets, + ) + }, + ) -def parse_stereo_observations( - left_image_files: List[Path], - right_image_files: List[Path], - joint_states_files: List[Path], - left_target_files: List[Path], - right_target_files: List[Path], -) -> StereoObservations: - r"""Parse stereo observations. - - Args: - left_image_files (List[Path]): Left image files. - right_image_files (List[Path]): Right image files. - joint_states_files (List[Path]): Joint states files. - left_target_files (List[Path]): Left target files. - right_target_files (List[Path]): Right target files. - Returns: - StereoObservations: Data for stereo registration. - """ - if ( - len(left_image_files) != len(right_image_files) - or len(left_image_files) != len(joint_states_files) - or len(left_image_files) != len(left_target_files) - or len(left_image_files) != len(right_target_files) - ): +def parse_stereo_observations( + left_image_files: list[Path] | None, + right_image_files: list[Path] | None, + joint_states_files: list[Path], + left_target_files: list[Path], + right_target_files: list[Path], +) -> ImageObservations: + r"""Parse stereo image-registration observations.""" + + lengths = { + "joint_states": len(joint_states_files), + "left_targets": len(left_target_files), + "right_targets": len(right_target_files), + } + + if left_image_files is not None: + lengths["left_images"] = len(left_image_files) + + if right_image_files is not None: + lengths["right_images"] = len(right_image_files) + + if len(set(lengths.values())) != 1: raise ValueError( - "Number of left / right images, joint states, left / right masks do not match." + f"All observation file lists must have the same length, got {lengths}." ) + if not joint_states_files: + raise ValueError("Expected at least one observation.") + rich.print("Parsing the following files:") - rich.print(f"Left images: {[f.name for f in left_image_files]}") - rich.print(f"Right images: {[f.name for f in right_image_files]}") - rich.print(f"Joint states: {[f.name for f in joint_states_files]}") - rich.print(f"Left targets: {[f.name for f in left_target_files]}") - rich.print(f"Right targets: {[f.name for f in right_target_files]}") + if left_image_files is not None: + rich.print(f"Left images: {[path.name for path in left_image_files]}") + if right_image_files is not None: + rich.print(f"Right images: {[path.name for path in right_image_files]}") + rich.print(f"Joint states: {[path.name for path in joint_states_files]}") + rich.print(f"Left targets: {[path.name for path in left_target_files]}") + rich.print(f"Right targets: {[path.name for path in right_target_files]}") + + left_images = ( + [_read_image(path) for path in left_image_files] + if left_image_files is not None + else None + ) + right_images = ( + [_read_image(path) for path in right_image_files] + if right_image_files is not None + else None + ) - left_images = [cv2.imread(f) for f in left_image_files] - right_images = [cv2.imread(f) for f in right_image_files] - joint_states = [np.load(f) for f in joint_states_files] - left_targets = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in left_target_files] - right_targets = [cv2.imread(f, cv2.IMREAD_GRAYSCALE) for f in right_target_files] - return StereoObservations( - left_images=left_images, - right_images=right_images, + joint_states = [np.load(path) for path in joint_states_files] + left_targets = [_read_target(path) for path in left_target_files] + right_targets = [_read_target(path) for path in right_target_files] + + if left_images is not None: + _validate_image_target_shapes( + images=left_images, + targets=left_targets, + camera_name="left", + ) + + if right_images is not None: + _validate_image_target_shapes( + images=right_images, + targets=right_targets, + camera_name="right", + ) + + return ImageObservations( joint_states=joint_states, - left_targets=left_targets, - right_targets=right_targets, + cameras={ + "left": CameraObservations( + images=left_images, + targets=left_targets, + ), + "right": CameraObservations( + images=right_images, + targets=right_targets, + ), + }, ) diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py index 7af110e..a9b1c8c 100644 --- a/roboreg/registration/image/config.py +++ b/roboreg/registration/image/config.py @@ -54,7 +54,7 @@ def __post_init__(self) -> None: @dataclass(frozen=True) -class DRRegConfig: +class DiffRenderingRegistrationConfig: camera: CameraConfig = field(default_factory=CameraConfig) optimizer: str = "AdamW" @@ -71,7 +71,7 @@ def __post_init__(self) -> None: @dataclass(frozen=True) -class CSRegConfig: +class CameraSwarmRegistrationConfig: camera: CameraConfig = field(default_factory=CameraConfig) n_cameras: int = 50 diff --git a/roboreg/registration/image/request.py b/roboreg/registration/image/request.py index 832c336..21ccb96 100644 --- a/roboreg/registration/image/request.py +++ b/roboreg/registration/image/request.py @@ -1,5 +1,4 @@ from dataclasses import dataclass -from typing import List import numpy as np @@ -13,84 +12,69 @@ @dataclass(frozen=True) -class MonocularObservations: - images: List[np.ndarray] - joint_states: List[np.ndarray] - targets: List[np.ndarray] +class CameraData: + intrinsics: np.ndarray + extrinsics: np.ndarray def __post_init__(self) -> None: - lengths = { - "images": len(self.images), - "joint_states": len(self.joint_states), - "targets": len(self.targets), - } + validate_intrinsics(self.intrinsics) + validate_extrinsics(self.extrinsics) + - if len(set(lengths.values())) != 1: +@dataclass(frozen=True) +class CameraObservations: + targets: list[np.ndarray] + images: list[np.ndarray] | None = None + + def __post_init__(self) -> None: + if not self.targets: + raise ValueError("Expected at least one target.") + + if self.images is not None and len(self.images) != len(self.targets): raise ValueError( - f"All observation fields must have the same length, got {lengths}." + "Expected the same number of images and targets, " + f"got {len(self.images)} and {len(self.targets)}." ) - if not self.images: - raise ValueError("Expected at least one observation.") - - validate_images(self.images, "images") validate_targets(self.targets, "targets") + target_shape = self.targets[0].shape[:2] + if any(target.shape[:2] != target_shape for target in self.targets): + raise ValueError("Expected all targets to have the same shape.") -@dataclass(frozen=True) -class StereoObservations: - left_images: List[np.ndarray] - right_images: List[np.ndarray] - joint_states: List[np.ndarray] - left_targets: List[np.ndarray] - right_targets: List[np.ndarray] + if self.images is not None: + validate_images(self.images, "images") - def __post_init__(self) -> None: - lengths = { - "left_images": len(self.left_images), - "right_images": len(self.right_images), - "joint_states": len(self.joint_states), - "left_targets": len(self.left_targets), - "right_targets": len(self.right_targets), - } - - if len(set(lengths.values())) != 1: - raise ValueError( - f"All observation fields must have the same length, got {lengths}." - ) + image_shape = self.images[0].shape[:2] + if any(image.shape[:2] != image_shape for image in self.images): + raise ValueError("Expected all images to have the same shape.") - if not self.left_images: - raise ValueError("Expected at least one observation.") + if image_shape != target_shape: + raise ValueError( + f"Image shape {image_shape} does not match " + f"target shape {target_shape}." + ) - validate_images(self.left_images, "left_images") - validate_images(self.right_images, "right_images") - validate_targets(self.left_targets, "left_targets") - validate_targets(self.right_targets, "right_targets") + @property + def shape(self) -> tuple[int, int]: + return self.targets[0].shape[:2] @dataclass(frozen=True) -class MonocularRequest: - intrinsics: np.ndarray - robot_data: RobotData - observations: MonocularObservations - initial_extrinsics: np.ndarray - - def __post_init__(self) -> None: - validate_intrinsics(self.intrinsics) - validate_extrinsics(self.initial_extrinsics) +class ImageObservations: + joint_states: list[np.ndarray] + cameras: dict[str, CameraObservations] @dataclass(frozen=True) -class StereoRequest: - left_intrinsics: np.ndarray - right_intrinsics: np.ndarray - initial_left_extrinsics: np.ndarray - left_to_right_extrinsics: np.ndarray +class ImageRegistrationRequest: + cameras: dict[str, CameraData] robot_data: RobotData - observations: StereoObservations + observations: ImageObservations + initial_extrinsics: np.ndarray def __post_init__(self) -> None: - validate_intrinsics(self.left_intrinsics) - validate_intrinsics(self.right_intrinsics) - validate_extrinsics(self.initial_left_extrinsics) - validate_extrinsics(self.left_to_right_extrinsics) + if not self.cameras: + raise ValueError("Expected at least one camera.") + + validate_extrinsics(self.initial_extrinsics) diff --git a/roboreg/registration/image/solver.py b/roboreg/registration/image/solver.py index 0c04e39..56deee7 100644 --- a/roboreg/registration/image/solver.py +++ b/roboreg/registration/image/solver.py @@ -1,15 +1,23 @@ import torch -from roboreg.registration.image.config import CSRegConfig, DRRegConfig +from roboreg.core.robot import Robot +from roboreg.core.scene import RobotScene +from roboreg.core.structs import VirtualCamera +from roboreg.registration.image.config import ( + CameraSwarmRegistrationConfig, + DiffRenderingRegistrationConfig, +) from roboreg.registration.image.objectives import RenderingObjective -from roboreg.registration.image.request import MonocularRequest, StereoRequest +from roboreg.registration.image.request import ImageRegistrationRequest from roboreg.registration.result import RegistrationResult +from roboreg.core.rendering import NVDiffRastRenderer +from roboreg.util.transform import rescale_intrinsics -class MonocularDiffRendRegistration: +class DiffRenderingRegistration: def __init__( self, - config: DRRegConfig, + config: DiffRenderingRegistrationConfig, objective: RenderingObjective, device: torch.device | str = "cuda", ) -> None: @@ -19,30 +27,55 @@ def __init__( def __call__( self, - request: MonocularRequest, + request: ImageRegistrationRequest, ) -> RegistrationResult: + robot_scene = self._prepare_robot_scene(request) pass - -class StereoDiffRendRegistration: - def __init__( + def _prepare_robot_scene( self, - config: DRRegConfig, - objective: RenderingObjective, - device: torch.device | str = "cuda", - ) -> None: - self._config = config - self._objective = objective - self._device = torch.device(device) - - def __call__(self, request: StereoRequest) -> RegistrationResult: - pass + request: ImageRegistrationRequest, + ) -> RobotScene: + # prepare cameras + cameras: dict[str, VirtualCamera] = {} + for camera_name, camera_data in request.cameras.items(): + # handle intrinsic scaling between hardware resolution and rendering resolution + camera_observations = request.observations.cameras[camera_name] + native_resolution = camera_observations.shape + target_resolution = ( + self._config.camera.target_resolution or native_resolution + ) + intrinsics = rescale_intrinsics( + intrinsics=camera_data.intrinsics, + source_resolution=native_resolution, + target_resolution=target_resolution, + ) + # prepare virtual camera + cameras[camera_name] = VirtualCamera( + resolution=target_resolution, + intrinsics=intrinsics, + extrinsics=camera_data.extrinsics, + z_min=self._config.camera.z_min, + z_max=self._config.camera.z_max, + device=self._device, + ) + # prepare robot + robot = Robot.from_robot_data( + robot_data=request.robot_data, + batch_size=len(request.observations.joint_states), + device=self._device, + ) + return RobotScene( + cameras=cameras, + robot=robot, + renderer=NVDiffRastRenderer(device=self._device), + ) class CameraSwarmRegistration: def __init__( self, - config: CSRegConfig, + config: CameraSwarmRegistrationConfig, objective: RenderingObjective, device: torch.device | str = "cuda", ) -> None: @@ -50,5 +83,5 @@ def __init__( self._objective = objective self._device = torch.device(device) - def __call__(self, request: MonocularRequest) -> RegistrationResult: + def __call__(self, request: ImageRegistrationRequest) -> RegistrationResult: pass diff --git a/roboreg/util/transform.py b/roboreg/util/transform.py index 4362c5c..5b6a264 100644 --- a/roboreg/util/transform.py +++ b/roboreg/util/transform.py @@ -1,5 +1,6 @@ -from typing import Optional +from typing import Optional, Tuple, Union +import numpy as np import torch @@ -128,3 +129,20 @@ def look_at_from_angle( random_rot[:, 1, 1] = torch.cos(angle) return random_ht @ random_rot + + +def rescale_intrinsics( + intrinsics: Union[np.ndarray, torch.Tensor], + current_resolution: Tuple[int, int], + target_resolution: Tuple[int, int], +) -> Union[np.ndarray, torch.Tensor]: + scaled = ( + intrinsics.copy() if isinstance(intrinsics, np.ndarray) else intrinsics.clone() + ) + scale_x = target_resolution[1] / current_resolution[1] + scale_y = target_resolution[0] / current_resolution[0] + scaled[..., 0, 0] *= scale_x + scaled[..., 1, 1] *= scale_y + scaled[..., 0, 2] *= scale_x + scaled[..., 1, 2] *= scale_y + return scaled diff --git a/test/io/test_parsers.py b/test/io/test_parsers.py index f813fec..09ff687 100644 --- a/test/io/test_parsers.py +++ b/test/io/test_parsers.py @@ -128,22 +128,29 @@ def test_parse_monocular_observations() -> None: ) assert ( - len(observations.images) + len(observations.cameras["camera"].images) == len(observations.joint_states) - == len(observations.targets) + == len(observations.cameras["camera"].targets) ), "Expected same number of images / joint states / masks." - assert len(observations.images) >= 1, "Should at least have one sample." - assert observations.images[0].ndim == 3, "Expected 3D image (HxWx3)." - assert observations.images[0].shape[-1] == 3, "Expected 3 color channels." - assert observations.targets[0].ndim == 2, "Expected 2D mask." assert ( - observations.targets[0].dtype == np.uint8 + len(observations.cameras["camera"].images) >= 1 + ), "Should at least have one sample." + assert ( + observations.cameras["camera"].images[0].ndim == 3 + ), "Expected 3D image (HxWx3)." + assert ( + observations.cameras["camera"].images[0].shape[-1] == 3 + ), "Expected 3 color channels." + assert observations.cameras["camera"].targets[0].ndim == 2, "Expected 2D mask." + assert ( + observations.cameras["camera"].targets[0].dtype == np.uint8 ), "Expected unsigned integers for mask." - assert np.all(observations.targets[0] >= 0) and np.all( - observations.targets[0] <= 255 + assert np.all(observations.cameras["camera"].targets[0] >= 0) and np.all( + observations.cameras["camera"].targets[0] <= 255 ), "Expected mask in range [0, 255]." assert ( - observations.targets[0].shape[:2] == observations.images[0].shape[:2] + observations.cameras["camera"].targets[0].shape[:2] + == observations.cameras["camera"].images[0].shape[:2] ), "Mask and image dimensions should match." @@ -158,47 +165,54 @@ def test_parse_stereo_observations() -> None: ) assert ( - len(observations.left_images) - == len(observations.right_images) + len(observations.cameras["left"].images) + == len(observations.cameras["right"].images) == len(observations.joint_states) - == len(observations.left_targets) - == len(observations.right_targets) + == len(observations.cameras["left"].targets) + == len(observations.cameras["right"].targets) ), "Expected same number of left/right images, joint states, and left/right masks." - assert len(observations.left_images) >= 1, "Should at least have one sample." + assert ( + len(observations.cameras["left"].images) >= 1 + ), "Should at least have one sample." # Test left data - assert observations.left_images[0].ndim == 3, "Expected 3D left image (HxWx3)." assert ( - observations.left_images[0].shape[-1] == 3 + observations.cameras["left"].images[0].ndim == 3 + ), "Expected 3D left image (HxWx3)." + assert ( + observations.cameras["left"].images[0].shape[-1] == 3 ), "Expected 3 color channels for left image." - assert observations.left_targets[0].ndim == 2, "Expected 2D left mask." + assert observations.cameras["left"].targets[0].ndim == 2, "Expected 2D left mask." assert ( - observations.left_targets[0].dtype == np.uint8 + observations.cameras["left"].targets[0].dtype == np.uint8 ), "Expected unsigned integers for left mask." - assert np.all(observations.left_targets[0] >= 0) and np.all( - observations.left_targets[0] <= 255 + assert np.all(observations.cameras["left"].targets[0] >= 0) and np.all( + observations.cameras["left"].targets[0] <= 255 ), "Expected left mask in range [0, 255]." # Test right data - assert observations.right_images[0].ndim == 3, "Expected 3D right image (HxWx3)." assert ( - observations.right_images[0].shape[-1] == 3 + observations.cameras["right"].images[0].ndim == 3 + ), "Expected 3D right image (HxWx3)." + assert ( + observations.cameras["right"].images[0].shape[-1] == 3 ), "Expected 3 color channels for right image." - assert observations.right_targets[0].ndim == 2, "Expected 2D right mask." + assert observations.cameras["right"].targets[0].ndim == 2, "Expected 2D right mask." assert ( - observations.right_targets[0].dtype == np.uint8 + observations.cameras["right"].targets[0].dtype == np.uint8 ), "Expected unsigned integers for right mask." - assert np.all(observations.right_targets[0] >= 0) and np.all( - observations.right_targets[0] <= 255 + assert np.all(observations.cameras["right"].targets[0] >= 0) and np.all( + observations.cameras["right"].targets[0] <= 255 ), "Expected right mask in range [0, 255]." # Test dimensions match assert ( - observations.left_targets[0].shape[:2] == observations.left_images[0].shape[:2] + observations.cameras["left"].targets[0].shape[:2] + == observations.cameras["left"].images[0].shape[:2] ), "Left mask and image dimensions should match." assert ( - observations.right_targets[0].shape[:2] - == observations.right_images[0].shape[:2] + observations.cameras["right"].targets[0].shape[:2] + == observations.cameras["right"].images[0].shape[:2] ), "Right mask and image dimensions should match." diff --git a/test/util/test_transform.py b/test/util/test_transform.py index ceeee09..9d3f33d 100644 --- a/test/util/test_transform.py +++ b/test/util/test_transform.py @@ -16,6 +16,7 @@ from_homogeneous, generate_ht_optical, look_at_from_angle, + rescale_intrinsics, to_homogeneous, ) @@ -152,7 +153,55 @@ def test_look_at_from_angle() -> None: raise ValueError(f"Expected shape ({batch_size}, 4, 4), got {ht.shape}.") +@pytest.mark.parametrize("backend", ["numpy", "torch"]) +def test_rescale_intrinsics(backend: str) -> None: + intrinsics_np = np.array( + [ + [100.0, 0.0, 40.0], + [0.0, 200.0, 30.0], + [0.0, 0.0, 1.0], + ], + dtype=np.float32, + ) + + # width x2, height x3 + resolution = (100, 100) + target_resolution = (300, 200) + + expected_np = np.array( + [ + [200.0, 0.0, 80.0], + [0.0, 600.0, 90.0], + [0.0, 0.0, 1.0], + ], + dtype=np.float32, + ) + + if backend == "numpy": + intrinsics = intrinsics_np.copy() + original = intrinsics.copy() + else: + intrinsics = torch.from_numpy(intrinsics_np.copy()) + original = intrinsics.clone() + + result = rescale_intrinsics( + intrinsics=intrinsics, + current_resolution=resolution, + target_resolution=target_resolution, + ) + + if backend == "numpy": + assert isinstance(result, np.ndarray) + np.testing.assert_allclose(result, expected_np) + np.testing.assert_array_equal(intrinsics, original) + else: + assert isinstance(result, torch.Tensor) + torch.testing.assert_close(result, torch.from_numpy(expected_np)) + torch.testing.assert_close(intrinsics, original) + + if __name__ == "__main__": test_depth_to_xyz() test_realsense_depth_to_xyz() test_look_at_from_angle() + test_rescale_intrinsics() From fb283f2864ad5293f346735cb7557a675ebec915 Mon Sep 17 00:00:00 2001 From: mhubii Date: Mon, 3 Aug 2026 15:09:43 +0100 Subject: [PATCH 10/12] functioning image-based refactor --- README.md | 22 +- cli/rr_cam_swarm.py | 10 +- cli/rr_hydra.py | 16 +- cli/rr_mono_dr.py | 250 ++++++----------- cli/rr_stereo_dr.py | 324 +++++++---------------- roboreg/io/parsers.py | 5 +- roboreg/registration/image/config.py | 32 +-- roboreg/registration/image/objectives.py | 23 ++ roboreg/registration/image/request.py | 5 +- roboreg/registration/image/solver.py | 154 ++++++++++- 10 files changed, 390 insertions(+), 451 deletions(-) diff --git a/README.md b/README.md index 22e5af3..af12263 100644 --- a/README.md +++ b/README.md @@ -183,9 +183,14 @@ This monocular differentiable rendering refinement requires a good initial estim ```shell rr-mono-dr \ - --optimizer SGD \ - --lr 0.01 \ - --max-iterations 100 \ + --optimizer AdamW \ + --lr 0.03 \ + --max-iterations 400 \ + --convergence-tolerance 0.001 \ + --convergence-patience 50 \ + --scheduler-factor 0.1 \ + --scheduler-patience 50 \ + --scheduler-threshold 0.0001 \ --display-progress \ --urdf-path test/assets/lbr_med7_r800/description/lbr_med7_r800.urdf \ --root-link-name lbr_link_0 \ @@ -209,9 +214,14 @@ This stereo differentiable rendering refinement requires a good initial estimate ```shell rr-stereo-dr \ - --optimizer SGD \ - --lr 0.01 \ - --max-iterations 100 \ + --optimizer AdamW \ + --lr 0.03 \ + --max-iterations 400 \ + --convergence-tolerance 0.001 \ + --convergence-patience 50 \ + --scheduler-factor 0.1 \ + --scheduler-patience 50 \ + --scheduler-threshold 0.0001 \ --display-progress \ --urdf-path test/assets/lbr_med7_r800/description/lbr_med7_r800.urdf \ --root-link-name lbr_link_0 \ diff --git a/cli/rr_cam_swarm.py b/cli/rr_cam_swarm.py index f8d373d..3b85c6c 100644 --- a/cli/rr_cam_swarm.py +++ b/cli/rr_cam_swarm.py @@ -1,5 +1,6 @@ import argparse import os +from pathlib import Path from typing import Union import cv2 @@ -252,14 +253,15 @@ def main() -> None: args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" os.environ["MAX_JOBS"] = str(args.max_jobs) # limit number of concurrent jobs + path = Path(args.path) # load data height, width, intrinsics = parse_camera_info( camera_info_file=args.camera_info_file ) - image_files = find_files(args.path, args.image_pattern) - target_files = find_files(args.path, args.mask_pattern) - joint_states_files = find_files(args.path, args.joint_states_pattern) + image_files = find_files(path, args.image_pattern) + target_files = find_files(path, args.mask_pattern) + joint_states_files = find_files(path, args.joint_states_pattern) n_samples = args.n_samples if n_samples > len(image_files): # randomly sample n_samples n_samples = len(image_files) @@ -418,7 +420,7 @@ def fitness_closure() -> torch.Tensor: HT_cam_swarm = look_at_from_angle( eye=best_eye, center=best_center, angle=best_angle ) - np.save(os.path.join(args.path, args.output_file), HT_cam_swarm.cpu().numpy()) + np.save(path / args.output_file), HT_cam_swarm.cpu().numpy() if __name__ == "__main__": diff --git a/cli/rr_hydra.py b/cli/rr_hydra.py index 7c93285..de40651 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -1,5 +1,5 @@ import argparse -import os +from pathlib import Path import numpy as np import torch @@ -187,19 +187,17 @@ def visualize_hydra_result( def main(): args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" + path = Path(args.path) # load data - joint_states_files = find_files(args.path, args.joint_states_pattern) - mask_files = find_files(args.path, args.mask_pattern) - depth_files = find_files(args.path, args.depth_pattern) observations = parse_hydra_observations( - joint_states_files=joint_states_files, - mask_files=mask_files, - depth_files=depth_files, + joint_states_files=find_files(path, args.joint_states_pattern), + mask_files=find_files(path, args.mask_pattern), + depth_files=find_files(path, args.depth_pattern), ) _, _, intrinsics = parse_camera_info(args.camera_info_file) - # instantiate robot + # load robot specifications if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -245,7 +243,7 @@ def main(): ) # to numpy - np.save(os.path.join(args.path, args.output_file), result.extrinsics.cpu().numpy()) + np.save(path / args.output_file, result.extrinsics.cpu().numpy()) if __name__ == "__main__": diff --git a/cli/rr_mono_dr.py b/cli/rr_mono_dr.py index 40064a8..ea99a5c 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -1,37 +1,33 @@ import argparse -import importlib import os -from enum import Enum +from pathlib import Path -import cv2 import numpy as np -import pytorch_kinematics as pk -import rich -import rich.progress import torch -from roboreg.core import NVDiffRastRenderer, Robot, RobotScene, VirtualCamera from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, + parse_camera_info, parse_monocular_observations, ) +from roboreg.registration.image.config import ( + CameraConfig, + ConvergenceConfig, + DiffRenderingRegistrationConfig, + PlateauSchedulerConfig, +) from roboreg.registration.image.objectives import ( - DistanceMapObjective, - ExponentialDecayMaskObjective, - RenderingObjective, + RenderingObjectiveType, + create_rendering_objective, ) -from roboreg.util import overlay_mask +from roboreg.registration.image.request import CameraData, ImageRegistrationRequest +from roboreg.registration.image.solver import DiffRenderingRegistration from .util.validate import validate_urdf_source -class REGISTRATION_MODE(Enum): - DISTANCE_FUNCTION = "distance-function" - SEGMENTATION = "segmentation" - - def args_factory() -> argparse.Namespace: parser = argparse.ArgumentParser( formatter_class=argparse.ArgumentDefaultsHelpFormatter @@ -39,39 +35,52 @@ def args_factory() -> argparse.Namespace: parser.add_argument( "--optimizer", type=str, - default="SGD", + default=DiffRenderingRegistrationConfig().optimizer, help="Optimizer to use, e.g. 'Adam' or 'SGD'. Imported from torch.optim.", ) parser.add_argument( "--lr", type=float, - default=1e-4, + default=DiffRenderingRegistrationConfig().lr, help="Learning rate for the optimizer.", ) parser.add_argument( "--max-iterations", type=int, - default=200, - help="Number of epochs to optimize for.", + default=ConvergenceConfig().max_iterations, + help="Maximum number of epochs to optimize for.", ) parser.add_argument( - "--step-size", + "--convergence-tolerance", + type=float, + default=ConvergenceConfig().tolerance, + ) + parser.add_argument( + "--convergence-patience", type=int, - default=100, - help="Step size for the learning rate scheduler.", + default=ConvergenceConfig().patience, ) parser.add_argument( - "--gamma", + "--scheduler-factor", type=float, - default=1.0, - help="Gamma for the learning rate scheduler.", + default=PlateauSchedulerConfig().factor, ) parser.add_argument( - "--mode", - type=str, - choices=[mode.value for mode in REGISTRATION_MODE], - default=REGISTRATION_MODE.DISTANCE_FUNCTION.value, - help="Registration mode.", + "--scheduler-patience", + type=int, + default=PlateauSchedulerConfig().patience, + ) + parser.add_argument( + "--scheduler-threshold", + type=float, + default=PlateauSchedulerConfig().threshold, + ) + parser.add_argument( + "--rendering-objective", + type=RenderingObjectiveType, + choices=list(RenderingObjectiveType), + default=RenderingObjectiveType.DISTANCE_MAP, + help="Rendering objective.", ) parser.add_argument( "--display-progress", @@ -167,46 +176,18 @@ def main() -> None: args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" os.environ["MAX_JOBS"] = str(args.max_jobs) # limit number of concurrent jobs - mode = REGISTRATION_MODE(args.mode) + path = Path(args.path) # load data - image_files = find_files(args.path, args.image_pattern) - joint_states_files = find_files(args.path, args.joint_states_pattern) - target_files = find_files(args.path, args.mask_pattern) observations = parse_monocular_observations( - image_files=image_files, - joint_states_files=joint_states_files, - target_files=target_files, - ) - - # pre-process data - joint_states = torch.tensor( - np.stack(observations.joint_states, axis=0), dtype=torch.float32, device=device - ) - objective: RenderingObjective - if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - objective = DistanceMapObjective() - elif mode == REGISTRATION_MODE.SEGMENTATION: - objective = ExponentialDecayMaskObjective() - else: - raise ValueError("Invalid registration mode.") - preprocessed_targets = objective.preprocess_targets( - targets=torch.tensor( - np.stack(observations.cameras["camera"].targets).astype(np.float32) / 255.0, - dtype=torch.float32, - device=device, - ) + image_files=find_files(path, args.image_pattern), + joint_states_files=find_files(path, args.joint_states_pattern), + target_files=find_files(path, args.mask_pattern), ) + _, _, intrinsics = parse_camera_info(args.camera_info_file) + extrinsics = np.load(args.extrinsics_file) - # instantiate camera with default identity extrinsics because we optimize for robot pose instead - camera = { - "camera": VirtualCamera.from_camera_configs( - camera_info_file=args.camera_info_file, - device=device, - ) - } - - # instantiate robot + # load robot specifications if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -222,114 +203,45 @@ def main() -> None: end_link_name=args.end_link_name, collision=args.collision_meshes, ) - robot = Robot.from_robot_data( - robot_data=robot_data, batch_size=joint_states.shape[0], device=device - ) - - # instantiate scene - scene = RobotScene( - cameras=camera, - robot=robot, - renderer=NVDiffRastRenderer(device=device), - ) - # load extrinsics estimate - extrinsics = torch.tensor( - np.load(args.extrinsics_file), dtype=torch.float32, device=device - ) - extrinsics_inv = torch.linalg.inv(extrinsics) - - # enable gradient tracking and instantiate optimizer - extrinsics_9d_inv = pk.matrix44_to_se3_9d(extrinsics_inv) - extrinsics_9d_inv.requires_grad = True - optimizer = getattr(importlib.import_module("torch.optim"), args.optimizer)( - [extrinsics_9d_inv], lr=args.lr - ) - scheduler = torch.optim.lr_scheduler.StepLR( - optimizer, step_size=args.step_size, gamma=args.gamma - ) - best_extrinsics = extrinsics - best_extrinsics_inv = extrinsics_inv - best_loss = float("inf") - - for iteration in rich.progress.track( - range(1, args.max_iterations + 1), "Optimizing..." - ): - if not extrinsics_9d_inv.requires_grad: - raise ValueError("Extrinsics require gradients.") - if not torch.is_grad_enabled(): - raise ValueError("Gradients must be enabled.") - extrinsics_inv = pk.se3_9d_to_matrix44(extrinsics_9d_inv) - scene.robot.configure(joint_states, extrinsics_inv) - renders = { - "camera": scene.observe_from("camera"), - } - loss = objective( - preprocessed_targets=preprocessed_targets, renders=renders["camera"] - ) - optimizer.zero_grad() - loss.backward() - optimizer.step() - scheduler.step() - - rich.print( - f"Step [{iteration} / {args.max_iterations}], loss: {np.round(loss.item(), 3)}, best loss: {np.round(best_loss, 3)}, lr: {scheduler.get_last_lr().pop()}" - ) - - if loss.item() < best_loss: - best_loss = loss.item() - best_extrinsics_inv = extrinsics_inv.detach().clone() - best_extrinsics = torch.linalg.inv(best_extrinsics_inv) - - # display optimization progress - if args.display_progress: - render = renders["camera"][0].squeeze().detach().cpu().numpy() - image = observations.cameras["camera"].images[0] - render_overlay = overlay_mask( - image, - (render * 255.0).astype(np.uint8), - scale=1.0, - ) - # difference left / right render / mask - difference = ( - cv2.cvtColor( - np.abs( - render - - observations.cameras["camera"].targets[0].astype(np.float32) - / 255.0 - ), - cv2.COLOR_GRAY2BGR, + # register + diff_rendering_registration = DiffRenderingRegistration( + config=DiffRenderingRegistrationConfig( + camera=CameraConfig(), + optimizer=args.optimizer, + lr=args.lr, + convergence=ConvergenceConfig( + max_iterations=args.max_iterations, + tolerance=args.convergence_tolerance, + patience=args.convergence_patience, + ), + plateau_scheduler=PlateauSchedulerConfig( + mode="min", + factor=args.scheduler_factor, + patience=args.scheduler_patience, + threshold=args.scheduler_threshold, + ), + ), + objective=create_rendering_objective(objective_type=args.rendering_objective), + device=device, + ) + result = diff_rendering_registration( + request=ImageRegistrationRequest( + cameras={ + "camera": CameraData( + intrinsics=intrinsics, ) - * 255.0 - ).astype(np.uint8) - # overlay segmentation mask - segmentation_overlay = overlay_mask( - image, - observations.cameras["camera"].targets[0], - mode="b", - scale=1.0, - ) - cv2.imshow( - "left to right: render overlay, difference, segmentation overlay", - cv2.resize( - np.hstack( - [ - render_overlay, - difference, - segmentation_overlay, - ] - ), - (0, 0), - fx=0.5, - fy=0.5, - ), - ) - cv2.waitKey(1) + }, + robot_data=robot_data, + observations=observations, + initial_extrinsics=extrinsics, + ) + ) # save extrinsics np.save( - os.path.join(args.path, args.output_file), - best_extrinsics.cpu().numpy(), + path / args.output_file, + result.extrinsics.cpu().numpy(), ) diff --git a/cli/rr_stereo_dr.py b/cli/rr_stereo_dr.py index b0bf448..964f87e 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -1,37 +1,33 @@ import argparse -import importlib import os -from enum import Enum +from pathlib import Path -import cv2 import numpy as np -import pytorch_kinematics as pk -import rich -import rich.progress import torch -from roboreg.core import NVDiffRastRenderer, Robot, RobotScene, VirtualCamera from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, + parse_camera_info, parse_stereo_observations, ) +from roboreg.registration.image.config import ( + CameraConfig, + ConvergenceConfig, + DiffRenderingRegistrationConfig, + PlateauSchedulerConfig, +) from roboreg.registration.image.objectives import ( - DistanceMapObjective, - ExponentialDecayMaskObjective, - RenderingObjective, + RenderingObjectiveType, + create_rendering_objective, ) -from roboreg.util import overlay_mask +from roboreg.registration.image.request import CameraData, ImageRegistrationRequest +from roboreg.registration.image.solver import DiffRenderingRegistration from .util.validate import validate_urdf_source -class REGISTRATION_MODE(Enum): - DISTANCE_FUNCTION = "distance-function" - SEGMENTATION = "segmentation" - - def args_factory() -> argparse.Namespace: parser = argparse.ArgumentParser( formatter_class=argparse.ArgumentDefaultsHelpFormatter @@ -39,39 +35,52 @@ def args_factory() -> argparse.Namespace: parser.add_argument( "--optimizer", type=str, - default="SGD", + default=DiffRenderingRegistrationConfig().optimizer, help="Optimizer to use, e.g. 'Adam' or 'SGD'. Imported from torch.optim.", ) parser.add_argument( "--lr", type=float, - default=1e-4, + default=DiffRenderingRegistrationConfig().lr, help="Learning rate for the optimizer.", ) parser.add_argument( "--max-iterations", type=int, - default=200, - help="Number of epochs to optimize for.", + default=ConvergenceConfig().max_iterations, + help="Maximum number of epochs to optimize for.", ) parser.add_argument( - "--step-size", + "--convergence-tolerance", + type=float, + default=ConvergenceConfig().tolerance, + ) + parser.add_argument( + "--convergence-patience", type=int, - default=100, - help="Step size for the learning rate scheduler.", + default=ConvergenceConfig().patience, ) parser.add_argument( - "--gamma", + "--scheduler-factor", type=float, - default=1.0, - help="Gamma for the learning rate scheduler.", + default=PlateauSchedulerConfig().factor, ) parser.add_argument( - "--mode", - type=str, - choices=[mode.value for mode in REGISTRATION_MODE], - default=REGISTRATION_MODE.DISTANCE_FUNCTION.value, - help="Registration mode.", + "--scheduler-patience", + type=int, + default=PlateauSchedulerConfig().patience, + ) + parser.add_argument( + "--scheduler-threshold", + type=float, + default=PlateauSchedulerConfig().threshold, + ) + parser.add_argument( + "--rendering-objective", + type=RenderingObjectiveType, + choices=list(RenderingObjectiveType), + default=RenderingObjectiveType.DISTANCE_MAP, + help="Rendering objective.", ) parser.add_argument( "--display-progress", @@ -197,64 +206,23 @@ def main() -> None: args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" os.environ["MAX_JOBS"] = str(args.max_jobs) # limit number of concurrent jobs - mode = REGISTRATION_MODE(args.mode) + path = Path(args.path) # load data - left_image_files = find_files(args.path, args.left_image_pattern) - right_image_files = find_files(args.path, args.right_image_pattern) - joint_states_files = find_files(args.path, args.joint_states_pattern) - left_target_files = find_files(args.path, args.left_mask_pattern) - right_target_files = find_files(args.path, args.right_mask_pattern) observations = parse_stereo_observations( - left_image_files=left_image_files, - right_image_files=right_image_files, - joint_states_files=joint_states_files, - left_target_files=left_target_files, - right_target_files=right_target_files, - ) - - # pre-process data - joint_states = torch.tensor( - np.stack(observations.joint_states, axis=0), dtype=torch.float32, device=device - ) - objective: RenderingObjective - if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - objective = DistanceMapObjective() - elif mode == REGISTRATION_MODE.SEGMENTATION: - objective = ExponentialDecayMaskObjective() - else: - raise ValueError("Invalid registration mode.") - left_preprocessed_targets = objective.preprocess_targets( - targets=torch.tensor( - np.stack(observations.cameras["left"].targets).astype(np.float32) / 255.0, - dtype=torch.float32, - device=device, - ) - ) - right_preprocessed_targets = objective.preprocess_targets( - targets=torch.tensor( - np.stack(observations.cameras["right"].targets).astype(np.float32) / 255.0, - dtype=torch.float32, - device=device, - ) + left_image_files=find_files(path, args.left_image_pattern), + right_image_files=find_files(path, args.right_image_pattern), + joint_states_files=find_files(path, args.joint_states_pattern), + left_target_files=find_files(path, args.left_mask_pattern), + right_target_files=find_files(path, args.right_mask_pattern), ) - # instantiate: - # - left camera with default identity extrinsics because we optimize for robot pose instead - # - right camera with transformation to left camera frame - cameras = { - "left": VirtualCamera.from_camera_configs( - camera_info_file=args.left_camera_info_file, - device=device, - ), - "right": VirtualCamera.from_camera_configs( - camera_info_file=args.right_camera_info_file, - extrinsics_file=args.right_extrinsics_file, - device=device, - ), - } + _, _, left_intrinsics = parse_camera_info(args.left_camera_info_file) + _, _, right_intrinsics = parse_camera_info(args.right_camera_info_file) + extrinsics = np.load(args.left_extrinsics_file) + right_extrinsics = np.load(args.right_extrinsics_file) - # instantiate robot + # load robot specifications if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -270,163 +238,53 @@ def main() -> None: end_link_name=args.end_link_name, collision=args.collision_meshes, ) - robot = Robot.from_robot_data( - robot_data=robot_data, batch_size=joint_states.shape[0], device=device - ) - - # instantiate scene - scene = RobotScene( - cameras=cameras, - robot=robot, - renderer=NVDiffRastRenderer(device=device), - ) - - # load extrinscis estimate...... - left_extrinsics = torch.tensor( - np.load(args.left_extrinsics_file), dtype=torch.float32, device=device - ) - left_extrinsics_inv = torch.linalg.inv(left_extrinsics) - - # enable gradient tracking and instantiate optimizer - left_extrinsics_9d_inv = pk.matrix44_to_se3_9d(left_extrinsics_inv) - left_extrinsics_9d_inv.requires_grad = True - optimizer = getattr(importlib.import_module("torch.optim"), args.optimizer)( - [left_extrinsics_9d_inv], lr=args.lr - ) - scheduler = torch.optim.lr_scheduler.StepLR( - optimizer, step_size=args.step_size, gamma=args.gamma - ) - best_left_extrinsics = left_extrinsics - best_left_extrinsics_inv = left_extrinsics_inv - best_loss = float("inf") - - for iteration in rich.progress.track( - range(1, args.max_iterations + 1), "Optimizing..." - ): - if not left_extrinsics_9d_inv.requires_grad: - raise ValueError("Extrinsics require gradients.") - if not torch.is_grad_enabled(): - raise ValueError("Gradients must be enabled.") - left_extrinsics_inv = pk.se3_9d_to_matrix44(left_extrinsics_9d_inv) - scene.robot.configure(joint_states, left_extrinsics_inv) - renders = { - "left": scene.observe_from("left"), - "right": scene.observe_from("right"), - } - loss = objective( - preprocessed_targets=left_preprocessed_targets, renders=renders["left"] - ) + objective( - preprocessed_targets=right_preprocessed_targets, renders=renders["right"] - ) - optimizer.zero_grad() - loss.backward() - optimizer.step() - scheduler.step() - - rich.print( - f"Step [{iteration} / {args.max_iterations}], loss: {np.round(loss.item(), 3)}, best loss: {np.round(best_loss, 3)}, lr: {scheduler.get_last_lr().pop()}" - ) - - if loss.item() < best_loss: - best_loss = loss.item() - best_left_extrinsics_inv = left_extrinsics_inv.detach().clone() - best_left_extrinsics = torch.linalg.inv(best_left_extrinsics_inv) - # display optimization progress - if args.display_progress: - render_overlays = [] - left_render = renders["left"][0].squeeze().detach().cpu().numpy() - left_image = observations.cameras["left"].images[0] - render_overlays.append( - overlay_mask( - left_image, - (left_render * 255.0).astype(np.uint8), - scale=1.0, - ) - ) - right_render = renders["right"][0].squeeze().detach().cpu().numpy() - right_image = observations.cameras["right"].images[0] - render_overlays.append( - overlay_mask( - right_image, - (right_render * 255.0).astype(np.uint8), - scale=1.0, - ) - ) - # difference left / right render / mask - differences = [] - differences.append( - ( - cv2.cvtColor( - np.abs( - left_render - - observations.cameras["left"].targets[0].astype(np.float32) - / 255.0 - ), - cv2.COLOR_GRAY2BGR, - ) - * 255.0 - ).astype(np.uint8) - ) - differences.append( - ( - cv2.cvtColor( - np.abs( - right_render - - observations.cameras["right"] - .targets[0] - .astype(np.float32) - / 255.0 - ), - cv2.COLOR_GRAY2BGR, - ) - * 255.0 - ).astype(np.uint8) - ) - # overlay segmentation mask - segmentation_overlays = [] - segmentation_overlays.append( - overlay_mask( - left_image, - observations.cameras["left"].targets[0], - mode="b", - scale=1.0, - ) - ) - segmentation_overlays.append( - overlay_mask( - right_image, - observations.cameras["right"].targets[0], - mode="b", - scale=1.0, - ) - ) - cv2.imshow( - "top to bottom: render overlays, differences, segmentation overlays | left: left view, right: right view", - cv2.resize( - np.vstack( - [ - np.hstack(render_overlays), - np.hstack(differences), - np.hstack(segmentation_overlays), - ] - ), - (0, 0), - fx=0.5, - fy=0.5, + # register + diff_rendering_registration = DiffRenderingRegistration( + config=DiffRenderingRegistrationConfig( + camera=CameraConfig(), + optimizer=args.optimizer, + lr=args.lr, + convergence=ConvergenceConfig( + max_iterations=args.max_iterations, + tolerance=args.convergence_tolerance, + patience=args.convergence_patience, + ), + plateau_scheduler=PlateauSchedulerConfig( + mode="min", + factor=args.scheduler_factor, + patience=args.scheduler_patience, + threshold=args.scheduler_threshold, + ), + ), + objective=create_rendering_objective(objective_type=args.rendering_objective), + device=device, + ) + result = diff_rendering_registration( + request=ImageRegistrationRequest( + cameras={ + "left": CameraData( + intrinsics=left_intrinsics, + ), + "right": CameraData( + intrinsics=right_intrinsics, + reference_to_camera=right_extrinsics, ), - ) - cv2.waitKey(1) + }, + robot_data=robot_data, + observations=observations, + initial_extrinsics=extrinsics, + ) + ) # save extrinsics np.save( - os.path.join(args.path, args.left_output_file), - best_left_extrinsics.cpu().numpy(), + path / args.left_output_file, + result.extrinsics.cpu().numpy(), ) np.save( - os.path.join(args.path, args.right_output_file), - best_left_extrinsics.cpu().numpy() - @ scene.cameras["right"].extrinsics.detach().cpu().numpy(), + path / args.right_output_file, + result.extrinsics.cpu().numpy() @ right_extrinsics, ) diff --git a/roboreg/io/parsers.py b/roboreg/io/parsers.py index b68651f..c66576f 100644 --- a/roboreg/io/parsers.py +++ b/roboreg/io/parsers.py @@ -391,7 +391,6 @@ def parse_monocular_observations( image_files: list[Path] | None, joint_states_files: list[Path], target_files: list[Path], - camera_name: str = "camera", ) -> ImageObservations: r"""Parse monocular image-registration observations.""" @@ -427,13 +426,13 @@ def parse_monocular_observations( _validate_image_target_shapes( images=images, targets=targets, - camera_name=camera_name, + camera_name="camera", ) return ImageObservations( joint_states=joint_states, cameras={ - camera_name: CameraObservations( + "camera": CameraObservations( images=images, targets=targets, ) diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py index a9b1c8c..7ced17d 100644 --- a/roboreg/registration/image/config.py +++ b/roboreg/registration/image/config.py @@ -23,34 +23,34 @@ def __post_init__(self) -> None: @dataclass(frozen=True) -class PlateauSchedulerConfig: - mode: Literal["min", "max"] = "min" - factor: float = 0.1 +class ConvergenceConfig: + max_iterations: int = 400 + tolerance: float = 1.0e-3 patience: int = 50 - threshold: float = 1.0e-4 def __post_init__(self) -> None: - if self.factor <= 0 or self.factor >= 1: - raise ValueError("factor must be in the range (0, 1).") + if self.max_iterations <= 0: + raise ValueError("max_iterations must be positive.") + if self.tolerance < 0: + raise ValueError("tolerance must be non-negative.") if self.patience < 0: raise ValueError("patience must be non-negative.") - if self.threshold < 0: - raise ValueError("threshold must be non-negative.") @dataclass(frozen=True) -class ConvergenceConfig: - max_iterations: int = 400 - tolerance: float = 1.0e-3 +class PlateauSchedulerConfig: + mode: Literal["min", "max"] = "min" + factor: float = 0.1 patience: int = 50 + threshold: float = 1.0e-4 def __post_init__(self) -> None: - if self.max_iterations <= 0: - raise ValueError("max_iterations must be positive.") - if self.tolerance < 0: - raise ValueError("tolerance must be non-negative.") + if self.factor <= 0 or self.factor >= 1: + raise ValueError("factor must be in the range (0, 1).") if self.patience < 0: raise ValueError("patience must be non-negative.") + if self.threshold < 0: + raise ValueError("threshold must be non-negative.") @dataclass(frozen=True) @@ -60,10 +60,10 @@ class DiffRenderingRegistrationConfig: optimizer: str = "AdamW" lr: float = 3.0e-2 + convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) plateau_scheduler: PlateauSchedulerConfig = field( default_factory=PlateauSchedulerConfig ) - convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) def __post_init__(self) -> None: if self.lr <= 0: diff --git a/roboreg/registration/image/objectives.py b/roboreg/registration/image/objectives.py index 0236829..86ec570 100644 --- a/roboreg/registration/image/objectives.py +++ b/roboreg/registration/image/objectives.py @@ -1,4 +1,5 @@ from dataclasses import dataclass +from enum import Enum from typing import Protocol import numpy as np @@ -161,3 +162,25 @@ def __call__( return soft_dice_loss( preprocessed_targets, renders, epsilon=self._config.epsilon ).mean() + + +class RenderingObjectiveType(str, Enum): + DISTANCE_MAP = "distance-map" + EXPONENTIAL_DECAY_MASK = "exponential-decay-mask" + PROBABILITY_MAP = "probability-map" + + def __str__(self) -> str: + return self.value + + +def create_rendering_objective( + objective_type: RenderingObjectiveType, +) -> RenderingObjective: + if objective_type == RenderingObjectiveType.DISTANCE_MAP: + return DistanceMapObjective() + elif objective_type == RenderingObjectiveType.EXPONENTIAL_DECAY_MASK: + return ExponentialDecayMaskObjective() + elif objective_type == RenderingObjectiveType.PROBABILITY_MAP: + return ProbabilityMapObjective() + else: + raise ValueError(f"Unsupported objective type: {objective_type}") diff --git a/roboreg/registration/image/request.py b/roboreg/registration/image/request.py index 21ccb96..facaf34 100644 --- a/roboreg/registration/image/request.py +++ b/roboreg/registration/image/request.py @@ -14,11 +14,12 @@ @dataclass(frozen=True) class CameraData: intrinsics: np.ndarray - extrinsics: np.ndarray + reference_to_camera: np.ndarray | None = None def __post_init__(self) -> None: validate_intrinsics(self.intrinsics) - validate_extrinsics(self.extrinsics) + if self.reference_to_camera is not None: + validate_extrinsics(self.reference_to_camera) @dataclass(frozen=True) diff --git a/roboreg/registration/image/solver.py b/roboreg/registration/image/solver.py index 56deee7..75c2bc5 100644 --- a/roboreg/registration/image/solver.py +++ b/roboreg/registration/image/solver.py @@ -1,5 +1,12 @@ +from typing import Iterable + +import numpy as np +import pytorch_kinematics as pk +import rich +import rich.progress import torch +from roboreg.core.rendering import NVDiffRastRenderer from roboreg.core.robot import Robot from roboreg.core.scene import RobotScene from roboreg.core.structs import VirtualCamera @@ -8,9 +15,11 @@ DiffRenderingRegistrationConfig, ) from roboreg.registration.image.objectives import RenderingObjective -from roboreg.registration.image.request import ImageRegistrationRequest -from roboreg.registration.result import RegistrationResult -from roboreg.core.rendering import NVDiffRastRenderer +from roboreg.registration.image.request import ( + ImageObservations, + ImageRegistrationRequest, +) +from roboreg.registration.result import RegistrationResult, TerminationReason from roboreg.util.transform import rescale_intrinsics @@ -25,21 +34,78 @@ def __init__( self._objective = objective self._device = torch.device(device) + # TODO: add support for callbacks + def __call__( self, request: ImageRegistrationRequest, ) -> RegistrationResult: - robot_scene = self._prepare_robot_scene(request) - pass + # PREPARE TARGETS (targets, joint states, extrinsics) + # prepare problem: enable gradient tracking and instantiate optimizer + extrinsics_inv = torch.linalg.inv( + torch.tensor( + request.initial_extrinsics, dtype=torch.float32, device=self._device + ) + ) + extrinsics_9d_inv = pk.matrix44_to_se3_9d(extrinsics_inv) + extrinsics_9d_inv.requires_grad = True + + if not extrinsics_9d_inv.requires_grad: + raise ValueError("Extrinsics require gradients.") + if not torch.is_grad_enabled(): + raise ValueError("Gradients must be enabled.") + + joint_states, preprocessed_targets = self._prepare_image_observations( + request.observations + ) + # PREPARE TARGETS END (targets, joint states, extrinsics) + robot_scene = self._create_robot_scene(request) + optimizer = self._create_optimizer(params=[extrinsics_9d_inv]) + scheduler = self._create_reduce_on_plateau_scheduler(optimizer) + # run optimization loop + result = self._optimize( + robot_scene=robot_scene, + joint_states=joint_states, + preprocessed_targets=preprocessed_targets, + extrinsics_9d_inv=extrinsics_9d_inv, + optimizer=optimizer, + scheduler=scheduler, + ) + return result + + def _prepare_image_observations( + self, + observations: ImageObservations, + ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: + joint_states = torch.as_tensor( + np.stack(observations.joint_states, axis=0), + dtype=torch.float32, + device=self._device, + ) + preprocessed_targets: dict[str, torch.Tensor] = {} + for camera_name, camera_observations in observations.cameras.items(): + targets = ( + torch.as_tensor( + np.stack(camera_observations.targets, axis=0), + dtype=torch.float32, + device=self._device, + ) + / 255.0 + ) + preprocessed_targets[camera_name] = self._objective.preprocess_targets( + targets + ) + return joint_states, preprocessed_targets - def _prepare_robot_scene( + def _create_robot_scene( self, request: ImageRegistrationRequest, ) -> RobotScene: # prepare cameras cameras: dict[str, VirtualCamera] = {} for camera_name, camera_data in request.cameras.items(): - # handle intrinsic scaling between hardware resolution and rendering resolution + # handle intrinsic scaling: in case of + # hardware resolution and rendering resolution mismatch camera_observations = request.observations.cameras[camera_name] native_resolution = camera_observations.shape target_resolution = ( @@ -47,14 +113,14 @@ def _prepare_robot_scene( ) intrinsics = rescale_intrinsics( intrinsics=camera_data.intrinsics, - source_resolution=native_resolution, + current_resolution=native_resolution, target_resolution=target_resolution, ) # prepare virtual camera cameras[camera_name] = VirtualCamera( resolution=target_resolution, intrinsics=intrinsics, - extrinsics=camera_data.extrinsics, + extrinsics=camera_data.reference_to_camera, z_min=self._config.camera.z_min, z_max=self._config.camera.z_max, device=self._device, @@ -71,6 +137,76 @@ def _prepare_robot_scene( renderer=NVDiffRastRenderer(device=self._device), ) + def _create_optimizer( + self, params: Iterable[torch.Tensor] + ) -> torch.optim.Optimizer: + return getattr(torch.optim, self._config.optimizer)(params, lr=self._config.lr) + + def _create_reduce_on_plateau_scheduler( + self, optimizer: torch.optim.Optimizer + ) -> torch.optim.lr_scheduler.ReduceLROnPlateau: + return torch.optim.lr_scheduler.ReduceLROnPlateau( + optimizer=optimizer, + mode=self._config.plateau_scheduler.mode, + factor=self._config.plateau_scheduler.factor, + patience=self._config.plateau_scheduler.patience, + threshold=self._config.plateau_scheduler.threshold, + ) + + def _optimize( + self, + robot_scene: RobotScene, + joint_states: torch.Tensor, + preprocessed_targets: dict[str, torch.Tensor], + extrinsics_9d_inv: torch.Tensor, + optimizer: torch.optim.Optimizer, + scheduler: torch.optim.lr_scheduler.ReduceLROnPlateau, + ) -> RegistrationResult: + + best_extrinsics_inv: torch.Tensor | None = None + best_loss = float("inf") + + # TODO: add convergence check... + + for iteration in rich.progress.track( + range(1, self._config.convergence.max_iterations + 1), "Optimizing..." + ): + + extrinsics_inv = pk.se3_9d_to_matrix44( + extrinsics_9d_inv + ) ### that's the parameter here.... + robot_scene.robot.configure(joint_states, extrinsics_inv) + + # per camera render and loss + camera_losses: list[torch.Tensor] = [] + for camera_name in robot_scene.cameras: + render = robot_scene.observe_from(camera_name) + camera_loss = self._objective( + preprocessed_targets=preprocessed_targets[camera_name], + renders=render, + ) + camera_losses.append(camera_loss) + loss = torch.stack(camera_losses).mean() + + optimizer.zero_grad() + loss.backward() + optimizer.step() + scheduler.step(metrics=loss) + + rich.print( + f"Step [{iteration} / {self._config.convergence.max_iterations}], loss: {np.round(loss.item(), 3)}, best loss: {np.round(best_loss, 3)}, lr: {scheduler.get_last_lr().pop()}" + ) + + if loss.item() < best_loss: + best_loss = loss.item() + best_extrinsics_inv = extrinsics_inv.detach().clone() + + return RegistrationResult( + extrinsics=torch.linalg.inv(best_extrinsics_inv), + iterations=iteration, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) + class CameraSwarmRegistration: def __init__( From 3aa145f2ac159be3be1be53acdec90d3a3fcd0cc Mon Sep 17 00:00:00 2001 From: mhubii Date: Tue, 4 Aug 2026 12:22:09 +0100 Subject: [PATCH 11/12] finished differentiable rendering refactor --- README.md | 8 +- cli/rr_hydra.py | 11 +- cli/rr_mono_dr.py | 37 +++++- cli/rr_stereo_dr.py | 37 +++++- roboreg/registration/image/callbacks.py | 40 ++++++ roboreg/registration/image/config.py | 4 +- roboreg/registration/image/solver.py | 138 ++++++++++++++------- roboreg/registration/point_cloud/solver.py | 16 +-- roboreg/registration/result.py | 3 + roboreg/util/viz.py | 9 +- 10 files changed, 237 insertions(+), 66 deletions(-) create mode 100644 roboreg/registration/image/callbacks.py diff --git a/README.md b/README.md index af12263..7d39d64 100644 --- a/README.md +++ b/README.md @@ -187,9 +187,9 @@ rr-mono-dr \ --lr 0.03 \ --max-iterations 400 \ --convergence-tolerance 0.001 \ - --convergence-patience 50 \ + --convergence-patience 100 \ --scheduler-factor 0.1 \ - --scheduler-patience 50 \ + --scheduler-patience 40 \ --scheduler-threshold 0.0001 \ --display-progress \ --urdf-path test/assets/lbr_med7_r800/description/lbr_med7_r800.urdf \ @@ -218,9 +218,9 @@ rr-stereo-dr \ --lr 0.03 \ --max-iterations 400 \ --convergence-tolerance 0.001 \ - --convergence-patience 50 \ + --convergence-patience 100 \ --scheduler-factor 0.1 \ - --scheduler-patience 50 \ + --scheduler-patience 40 \ --scheduler-threshold 0.0001 \ --display-progress \ --urdf-path test/assets/lbr_med7_r800/description/lbr_med7_r800.urdf \ diff --git a/cli/rr_hydra.py b/cli/rr_hydra.py index de40651..2775933 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -2,6 +2,7 @@ from pathlib import Path import numpy as np +import rich import torch from roboreg.io import ( @@ -232,8 +233,9 @@ def main(): hydra_robust_icp = HydraRobustICP( config=config, device=device, - callback=visualize_hydra_result if args.display_results else None, + on_after_registration=visualize_hydra_result if args.display_results else None, ) + rich.print("Entering optimization...") result = hydra_robust_icp( request=HydraRequest( intrinsics=intrinsics, @@ -241,8 +243,13 @@ def main(): observations=observations, ) ) + rich.print( + f"Optimization terminated after {result.iterations} iterations " + f"with status '{result.termination_reason}'." + ) - # to numpy + # save extrinsics + rich.print(f"Writing results to: '{path}'.") np.save(path / args.output_file, result.extrinsics.cpu().numpy()) diff --git a/cli/rr_mono_dr.py b/cli/rr_mono_dr.py index ea99a5c..82cbdf1 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -3,6 +3,7 @@ from pathlib import Path import numpy as np +import rich import torch from roboreg.io import ( @@ -12,6 +13,7 @@ parse_camera_info, parse_monocular_observations, ) +from roboreg.registration.image.callbacks import RenderOverlayCallback from roboreg.registration.image.config import ( CameraConfig, ConvergenceConfig, @@ -23,7 +25,11 @@ create_rendering_objective, ) from roboreg.registration.image.request import CameraData, ImageRegistrationRequest -from roboreg.registration.image.solver import DiffRenderingRegistration +from roboreg.registration.image.solver import ( + DiffRenderingRegistration, + OptimizationCallback, + OptimizationState, +) from .util.validate import validate_urdf_source @@ -172,6 +178,15 @@ def args_factory() -> argparse.Namespace: return parser.parse_args() +def print_optimization_state(state: OptimizationState) -> None: + rich.print( + f"Step [{state.iteration} / {state.max_iterations}], " + f"loss: {state.loss:.3f}, " + f"best loss: {state.best_loss:.3f}, " + f"lr: {state.learning_rate:.3e}" + ) + + def main() -> None: args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" @@ -205,6 +220,19 @@ def main() -> None: ) # register + on_iteration: list[OptimizationCallback] = [ + print_optimization_state, + ] + if args.display_progress: + on_iteration.append( + RenderOverlayCallback( + images={ + camera_name: camera_observations.images + for camera_name, camera_observations in observations.cameras.items() + if camera_observations.images is not None + }, + ) + ) diff_rendering_registration = DiffRenderingRegistration( config=DiffRenderingRegistrationConfig( camera=CameraConfig(), @@ -224,7 +252,9 @@ def main() -> None: ), objective=create_rendering_objective(objective_type=args.rendering_objective), device=device, + on_iteration=on_iteration, ) + rich.print("Entering optimization...") result = diff_rendering_registration( request=ImageRegistrationRequest( cameras={ @@ -237,8 +267,13 @@ def main() -> None: initial_extrinsics=extrinsics, ) ) + rich.print( + f"Optimization terminated after {result.iterations} iterations " + f"with status '{result.termination_reason}'." + ) # save extrinsics + rich.print(f"Writing results to: '{path}'.") np.save( path / args.output_file, result.extrinsics.cpu().numpy(), diff --git a/cli/rr_stereo_dr.py b/cli/rr_stereo_dr.py index 964f87e..21ae4ce 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -3,6 +3,7 @@ from pathlib import Path import numpy as np +import rich import torch from roboreg.io import ( @@ -12,6 +13,7 @@ parse_camera_info, parse_stereo_observations, ) +from roboreg.registration.image.callbacks import RenderOverlayCallback from roboreg.registration.image.config import ( CameraConfig, ConvergenceConfig, @@ -23,7 +25,11 @@ create_rendering_objective, ) from roboreg.registration.image.request import CameraData, ImageRegistrationRequest -from roboreg.registration.image.solver import DiffRenderingRegistration +from roboreg.registration.image.solver import ( + DiffRenderingRegistration, + OptimizationCallback, + OptimizationState, +) from .util.validate import validate_urdf_source @@ -202,6 +208,15 @@ def args_factory() -> argparse.Namespace: return parser.parse_args() +def print_optimization_state(state: OptimizationState) -> None: + rich.print( + f"Step [{state.iteration} / {state.max_iterations}], " + f"loss: {state.loss:.3f}, " + f"best loss: {state.best_loss:.3f}, " + f"lr: {state.learning_rate:.3e}" + ) + + def main() -> None: args = args_factory() device = "cuda" if torch.cuda.is_available() else "cpu" @@ -240,6 +255,19 @@ def main() -> None: ) # register + on_iteration: list[OptimizationCallback] = [ + print_optimization_state, + ] + if args.display_progress: + on_iteration.append( + RenderOverlayCallback( + images={ + camera_name: camera_observations.images + for camera_name, camera_observations in observations.cameras.items() + if camera_observations.images is not None + }, + ) + ) diff_rendering_registration = DiffRenderingRegistration( config=DiffRenderingRegistrationConfig( camera=CameraConfig(), @@ -259,7 +287,9 @@ def main() -> None: ), objective=create_rendering_objective(objective_type=args.rendering_objective), device=device, + on_iteration=[print_optimization_state], ) + rich.print("Entering optimization...") result = diff_rendering_registration( request=ImageRegistrationRequest( cameras={ @@ -276,8 +306,13 @@ def main() -> None: initial_extrinsics=extrinsics, ) ) + rich.print( + f"Optimization terminated after {result.iterations} iterations " + f"with status '{result.termination_reason}'." + ) # save extrinsics + rich.print(f"Writing results to: '{path}'.") np.save( path / args.left_output_file, result.extrinsics.cpu().numpy(), diff --git a/roboreg/registration/image/callbacks.py b/roboreg/registration/image/callbacks.py new file mode 100644 index 0000000..11f60d5 --- /dev/null +++ b/roboreg/registration/image/callbacks.py @@ -0,0 +1,40 @@ +import cv2 +import numpy as np + +from roboreg.registration.image.solver import OptimizationState +from roboreg.util.viz import overlay_mask + + +class RenderOverlayCallback: + def __init__( + self, + images: dict[str, list[np.ndarray]], + every_n_iterations: int = 1, + ) -> None: + self._images = images + self._every_n_iterations = every_n_iterations + + def __call__(self, state: OptimizationState) -> None: + if state.iteration % self._every_n_iterations != 0: + return + for camera_name, render in state.renders.items(): + images = self._images.get(camera_name) + if images is None: + continue + image = images[0] + mask = render[0].detach().cpu().numpy().squeeze() + mask = np.clip(mask * 255.0, 0, 255).astype(np.uint8) + if image.shape[:2] != mask.shape: + image = cv2.resize( + image, + (mask.shape[1], mask.shape[0]), + interpolation=cv2.INTER_LINEAR, + ) + overlay = overlay_mask( + img=image, + mask=mask, + mode="r", + scale=1.0, + ) + cv2.imshow(f"Render overlay: {camera_name}", overlay) + cv2.waitKey(1) diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py index 7ced17d..b658f92 100644 --- a/roboreg/registration/image/config.py +++ b/roboreg/registration/image/config.py @@ -26,7 +26,7 @@ def __post_init__(self) -> None: class ConvergenceConfig: max_iterations: int = 400 tolerance: float = 1.0e-3 - patience: int = 50 + patience: int = 100 def __post_init__(self) -> None: if self.max_iterations <= 0: @@ -41,7 +41,7 @@ def __post_init__(self) -> None: class PlateauSchedulerConfig: mode: Literal["min", "max"] = "min" factor: float = 0.1 - patience: int = 50 + patience: int = 40 threshold: float = 1.0e-4 def __post_init__(self) -> None: diff --git a/roboreg/registration/image/solver.py b/roboreg/registration/image/solver.py index 75c2bc5..592ca04 100644 --- a/roboreg/registration/image/solver.py +++ b/roboreg/registration/image/solver.py @@ -1,10 +1,10 @@ -from typing import Iterable +from dataclasses import dataclass +from typing import Callable, Iterable import numpy as np import pytorch_kinematics as pk -import rich -import rich.progress import torch +import torch.nn.functional as F from roboreg.core.rendering import NVDiffRastRenderer from roboreg.core.robot import Robot @@ -23,46 +23,45 @@ from roboreg.util.transform import rescale_intrinsics +@dataclass(frozen=True) +class OptimizationState: + iteration: int + max_iterations: int + loss: float + best_loss: float + learning_rate: float + extrinsics: torch.Tensor + renders: dict[str, torch.Tensor] + camera_losses: dict[str, float] + + +OptimizationCallback = Callable[[OptimizationState], None] + + class DiffRenderingRegistration: def __init__( self, config: DiffRenderingRegistrationConfig, objective: RenderingObjective, device: torch.device | str = "cuda", + on_iteration: list[OptimizationCallback] | None = None, ) -> None: self._config = config self._objective = objective self._device = torch.device(device) - - # TODO: add support for callbacks + self._on_iteration = on_iteration or [] def __call__( self, request: ImageRegistrationRequest, ) -> RegistrationResult: - # PREPARE TARGETS (targets, joint states, extrinsics) - # prepare problem: enable gradient tracking and instantiate optimizer - extrinsics_inv = torch.linalg.inv( - torch.tensor( - request.initial_extrinsics, dtype=torch.float32, device=self._device - ) - ) - extrinsics_9d_inv = pk.matrix44_to_se3_9d(extrinsics_inv) - extrinsics_9d_inv.requires_grad = True - - if not extrinsics_9d_inv.requires_grad: - raise ValueError("Extrinsics require gradients.") - if not torch.is_grad_enabled(): - raise ValueError("Gradients must be enabled.") - joint_states, preprocessed_targets = self._prepare_image_observations( request.observations ) - # PREPARE TARGETS END (targets, joint states, extrinsics) robot_scene = self._create_robot_scene(request) + extrinsics_9d_inv = self._prepare_extrinsics_9d_inv(request.initial_extrinsics) optimizer = self._create_optimizer(params=[extrinsics_9d_inv]) scheduler = self._create_reduce_on_plateau_scheduler(optimizer) - # run optimization loop result = self._optimize( robot_scene=robot_scene, joint_states=joint_states, @@ -92,6 +91,16 @@ def _prepare_image_observations( ) / 255.0 ) + # resize on hardware resolution, rendering resolution mismatch + target_resolution = ( + self._config.camera.target_resolution or camera_observations.shape + ) + if targets.shape[-2:] != target_resolution: + targets = F.interpolate( + targets.unsqueeze(1), + size=target_resolution, + mode="nearest", + ).squeeze(1) preprocessed_targets[camera_name] = self._objective.preprocess_targets( targets ) @@ -137,6 +146,22 @@ def _create_robot_scene( renderer=NVDiffRastRenderer(device=self._device), ) + def _prepare_extrinsics_9d_inv( + self, + initial_extrinsics: np.ndarray, + ) -> torch.Tensor: + extrinsics = torch.as_tensor( + initial_extrinsics, + dtype=torch.float32, + device=self._device, + ) + # TODO: Standardize transform naming and direction conventions + # https://github.com/lbr-stack/roboreg/issues/137 + extrinsics_inv = torch.linalg.inv(extrinsics) + return ( + pk.matrix44_to_se3_9d(extrinsics_inv).detach().clone().requires_grad_(True) + ) + def _create_optimizer( self, params: Iterable[torch.Tensor] ) -> torch.optim.Optimizer: @@ -162,45 +187,66 @@ def _optimize( optimizer: torch.optim.Optimizer, scheduler: torch.optim.lr_scheduler.ReduceLROnPlateau, ) -> RegistrationResult: - best_extrinsics_inv: torch.Tensor | None = None best_loss = float("inf") - - # TODO: add convergence check... - - for iteration in rich.progress.track( - range(1, self._config.convergence.max_iterations + 1), "Optimizing..." - ): - - extrinsics_inv = pk.se3_9d_to_matrix44( - extrinsics_9d_inv - ) ### that's the parameter here.... + iterations_without_improvement = 0 + for iteration in range(1, self._config.convergence.max_iterations + 1): + extrinsics_inv = pk.se3_9d_to_matrix44(extrinsics_9d_inv) robot_scene.robot.configure(joint_states, extrinsics_inv) - # per camera render and loss - camera_losses: list[torch.Tensor] = [] + camera_losses: dict[str, torch.Tensor] = {} + renders: dict[str, torch.Tensor] | None = {} if self._on_iteration else None for camera_name in robot_scene.cameras: render = robot_scene.observe_from(camera_name) - camera_loss = self._objective( + camera_losses[camera_name] = self._objective( preprocessed_targets=preprocessed_targets[camera_name], renders=render, ) - camera_losses.append(camera_loss) - loss = torch.stack(camera_losses).mean() - + if renders is not None: + renders[camera_name] = render + loss = torch.stack(list(camera_losses.values())).mean() optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step(metrics=loss) - - rich.print( - f"Step [{iteration} / {self._config.convergence.max_iterations}], loss: {np.round(loss.item(), 3)}, best loss: {np.round(best_loss, 3)}, lr: {scheduler.get_last_lr().pop()}" - ) - - if loss.item() < best_loss: - best_loss = loss.item() + loss_value = loss.item() + if loss_value < best_loss: + improvement = best_loss - loss_value + best_loss = loss_value best_extrinsics_inv = extrinsics_inv.detach().clone() + if improvement > self._config.convergence.tolerance: + iterations_without_improvement = 0 + else: + iterations_without_improvement += 1 + else: + iterations_without_improvement += 1 + if iterations_without_improvement >= self._config.convergence.patience: + return RegistrationResult( + extrinsics=torch.linalg.inv(best_extrinsics_inv), + iterations=iteration, + termination_reason=TerminationReason.CONVERGED, + ) + if self._on_iteration: + assert renders is not None + state = OptimizationState( + iteration=iteration, + max_iterations=self._config.convergence.max_iterations, + loss=loss_value, + best_loss=best_loss, + learning_rate=optimizer.param_groups[0]["lr"], + extrinsics=torch.linalg.inv(extrinsics_inv.detach()), + renders={ + camera_name: render.detach() + for camera_name, render in renders.items() + }, + camera_losses={ + camera_name: camera_loss.detach().item() + for camera_name, camera_loss in camera_losses.items() + }, + ) + for callback in self._on_iteration: + callback(state) return RegistrationResult( extrinsics=torch.linalg.inv(best_extrinsics_inv), iterations=iteration, diff --git a/roboreg/registration/point_cloud/solver.py b/roboreg/registration/point_cloud/solver.py index 25a095b..682f0bc 100644 --- a/roboreg/registration/point_cloud/solver.py +++ b/roboreg/registration/point_cloud/solver.py @@ -152,11 +152,11 @@ def __init__( self, config: HydraICPConfig | None = None, device: torch.device | str = "cuda", - callback: HydraCallback | None = None, + on_after_registration: HydraCallback | None = None, ) -> None: self._config = config or HydraICPConfig() self._device = torch.device(device) - self._callback = callback + self._on_after_registration = on_after_registration def __call__(self, request: HydraRequest) -> RegistrationResult: hydra_problem = _prepare_hydra_problem( @@ -176,8 +176,8 @@ def __call__(self, request: HydraRequest) -> RegistrationResult: max_iterations=self._config.max_iterations, rmse_change_tolerance=self._config.hydra.rmse_change_tolerance, ) - if self._callback is not None: - self._callback(hydra_problem, result) + if self._on_after_registration is not None: + self._on_after_registration(hydra_problem, result) return result @@ -186,11 +186,11 @@ def __init__( self, config: HydraRobustICPConfig | None = None, device: torch.device | str = "cuda", - callback: HydraCallback | None = None, + on_after_registration: HydraCallback | None = None, ) -> None: self._config = config or HydraRobustICPConfig() self._device = torch.device(device) - self._callback = callback + self._on_after_registration = on_after_registration def __call__(self, request: HydraRequest) -> RegistrationResult: hydra_problem = _prepare_hydra_problem( @@ -212,6 +212,6 @@ def __call__(self, request: HydraRequest) -> RegistrationResult: max_inner_iterations=self._config.max_inner_iterations, rmse_change_tolerance=self._config.hydra.rmse_change_tolerance, ) - if self._callback is not None: - self._callback(hydra_problem, result) + if self._on_after_registration is not None: + self._on_after_registration(hydra_problem, result) return result diff --git a/roboreg/registration/result.py b/roboreg/registration/result.py index f3b27df..c7cdbb4 100644 --- a/roboreg/registration/result.py +++ b/roboreg/registration/result.py @@ -9,6 +9,9 @@ class TerminationReason(str, Enum): MAX_ITERATIONS = "max_iterations" FAILED = "failed" + def __str__(self) -> str: + return self.value + @dataclass class RegistrationResult: diff --git a/roboreg/util/viz.py b/roboreg/util/viz.py index 36bfd98..aa2032a 100644 --- a/roboreg/util/viz.py +++ b/roboreg/util/viz.py @@ -1,4 +1,4 @@ -from typing import List, Optional +from typing import List, Literal, Optional import cv2 import numpy as np @@ -9,7 +9,7 @@ def overlay_mask( img: np.ndarray, mask: np.ndarray, - mode: str = "r", + mode: Literal["r", "g", "b"] = "r", alpha: float = 0.5, beta: float = 0.5, gamma: float = 0.0, @@ -27,6 +27,11 @@ def overlay_mask( Returns: Mask overlayed on image. """ + if img.shape[:2] != mask.shape: + raise ValueError( + f"Image and mask shapes must match, got " + f"{img.shape[:2]} and {mask.shape}." + ) colored_mask = None if mode == "r": colored_mask = np.stack( From f041a574c0c3108101bcd3f382a528ae18b28979 Mon Sep 17 00:00:00 2001 From: mhubii Date: Tue, 4 Aug 2026 12:32:08 +0100 Subject: [PATCH 12/12] remove camera swarm --- roboreg/registration/image/config.py | 29 ---------------------------- roboreg/registration/image/solver.py | 15 -------------- 2 files changed, 44 deletions(-) diff --git a/roboreg/registration/image/config.py b/roboreg/registration/image/config.py index b658f92..c63c494 100644 --- a/roboreg/registration/image/config.py +++ b/roboreg/registration/image/config.py @@ -68,32 +68,3 @@ class DiffRenderingRegistrationConfig: def __post_init__(self) -> None: if self.lr <= 0: raise ValueError("lr must be positive.") - - -@dataclass(frozen=True) -class CameraSwarmRegistrationConfig: - camera: CameraConfig = field(default_factory=CameraConfig) - - n_cameras: int = 50 - min_distance: float = 0.5 - max_distance: float = 2.0 - angle_range: float = math.pi - - inertia_weight: float = 0.7 - cognitive_coefficient: float = 1.5 - social_coefficient: float = 1.5 - - convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) - - def __post_init__(self) -> None: - if self.n_cameras <= 0: - raise ValueError("n_cameras must be positive.") - - if self.min_distance <= 0: - raise ValueError("min_distance must be positive.") - - if self.max_distance <= self.min_distance: - raise ValueError("max_distance must be greater than min_distance.") - - if self.angle_range <= 0: - raise ValueError("angle_range must be positive.") diff --git a/roboreg/registration/image/solver.py b/roboreg/registration/image/solver.py index 592ca04..4a36ee8 100644 --- a/roboreg/registration/image/solver.py +++ b/roboreg/registration/image/solver.py @@ -252,18 +252,3 @@ def _optimize( iterations=iteration, termination_reason=TerminationReason.MAX_ITERATIONS, ) - - -class CameraSwarmRegistration: - def __init__( - self, - config: CameraSwarmRegistrationConfig, - objective: RenderingObjective, - device: torch.device | str = "cuda", - ) -> None: - self._config = config - self._objective = objective - self._device = torch.device(device) - - def __call__(self, request: ImageRegistrationRequest) -> RegistrationResult: - pass