diff --git a/README.md b/README.md index 22e5af3..7d39d64 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 100 \ + --scheduler-factor 0.1 \ + --scheduler-patience 40 \ + --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 100 \ + --scheduler-factor 0.1 \ + --scheduler-patience 40 \ + --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 1a52e93..3b85c6c 100644 --- a/cli/rr_cam_swarm.py +++ b/cli/rr_cam_swarm.py @@ -1,25 +1,19 @@ import argparse import os +from pathlib import Path from typing import Union import cv2 import numpy as np 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, 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 @@ -259,33 +253,38 @@ 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) - mask_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) 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() - images, joint_states, masks = parse_mono_data( + observations = parse_monocular_observations( image_files=image_files, - mask_files=mask_files, + target_files=target_files, joint_states_files=joint_states_files, ) # pre-process data + camera_name = "camera" 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.cameras[camera_name].targets + ] masks = torch.tensor(np.array(masks), dtype=torch.float32, device=device) # scale image data (memory reduction) @@ -319,7 +318,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, @@ -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 @@ -380,10 +366,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) @@ -399,10 +385,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.cameras[camera_name].images[offset].shape[1], + observations.cameras[camera_name].images[offset].shape[0], + ), ) overlay = overlay_mask( - images[offset], + observations.cameras[camera_name].images[offset], current_best_render, scale=1.0, ) @@ -430,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 955e632..2775933 100644 --- a/cli/rr_hydra.py +++ b/cli/rr_hydra.py @@ -1,27 +1,25 @@ import argparse -import os +from pathlib import Path import numpy as np +import rich 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, load_robot_data_from_urdf_file, parse_camera_info, - parse_hydra_data, + 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 ( + DepthToPointCloudConfig, + HydraConfig, + HydraRobustICPConfig, ) +from roboreg.registration.point_cloud.request import HydraRequest +from roboreg.registration.point_cloud.solver import HydraProblem, HydraRobustICP +from roboreg.registration.result import RegistrationResult from .util.validate import validate_urdf_source @@ -167,23 +165,40 @@ 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" + 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) - joint_states, masks, depths = parse_hydra_data( - joint_states_files=joint_states_files, - mask_files=mask_files, - depth_files=depth_files, + observations = parse_hydra_observations( + 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), ) - height, width, intrinsics = parse_camera_info(args.camera_info_file) + _, _, intrinsics = parse_camera_info(args.camera_info_file) - # instantiate robot - batch_size = len(joint_states) + # load robot specifications if args.urdf_path is not None: robot_data = load_robot_data_from_urdf_file( urdf_path=args.urdf_path, @@ -199,118 +214,43 @@ def main(): end_link_name=args.end_link_name, collision=args.collision_meshes, ) - mesh_container = TorchMeshContainer( - meshes=robot_data.meshes, - batch_size=len(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(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) - 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] - ), - ) - ), + # register + config = HydraRobustICPConfig( + HydraConfig( + reference_points_per_mesh=args.number_of_points, + depth_to_point_cloud=DepthToPointCloudConfig( + 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, ) - for xyz, mask in zip(xyzs, 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 - - 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), + hydra_robust_icp = HydraRobustICP( + config=config, + device=device, + on_after_registration=visualize_hydra_result if args.display_results else None, + ) + rich.print("Entering optimization...") + result = hydra_robust_icp( + request=HydraRequest( + intrinsics=intrinsics, + robot_data=robot_data, + observations=observations, ) + ) + rich.print( + f"Optimization terminated after {result.iterations} iterations " + f"with status '{result.termination_reason}'." + ) - # to numpy - HT = HT.cpu().numpy() - np.save(os.path.join(args.path, args.output_file), HT) + # save extrinsics + rich.print(f"Writing results to: '{path}'.") + 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 9d6dc73..82cbdf1 100644 --- a/cli/rr_mono_dr.py +++ b/cli/rr_mono_dr.py @@ -1,40 +1,39 @@ 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, - TorchKinematics, - TorchMeshContainer, - VirtualCamera, -) from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, - parse_mono_data, + parse_camera_info, + parse_monocular_observations, +) +from roboreg.registration.image.callbacks import RenderOverlayCallback +from roboreg.registration.image.config import ( + CameraConfig, + ConvergenceConfig, + DiffRenderingRegistrationConfig, + PlateauSchedulerConfig, +) +from roboreg.registration.image.objectives import ( + RenderingObjectiveType, + create_rendering_objective, +) +from roboreg.registration.image.request import CameraData, ImageRegistrationRequest +from roboreg.registration.image.solver import ( + DiffRenderingRegistration, + OptimizationCallback, + OptimizationState, ) -from roboreg.losses import soft_dice_loss -from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask 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 @@ -42,39 +41,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", @@ -166,45 +178,31 @@ 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" 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) - mask_files = find_files(args.path, args.mask_pattern) - images, joint_states, masks = parse_mono_data( - image_files=image_files, - joint_states_files=joint_states_files, - mask_files=mask_files, + observations = parse_monocular_observations( + 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) - # pre-process data - joint_states = torch.tensor( - np.array(joint_states), dtype=torch.float32, device=device - ) - if mode == REGISTRATION_MODE.DISTANCE_FUNCTION: - targets = [mask_distance_transform(mask) for mask in masks] - elif mode == REGISTRATION_MODE.SEGMENTATION: - targets = [mask_exponential_decay(mask) for mask in masks] - else: - raise ValueError("Invalid registration mode.") - targets = torch.tensor( - np.array(targets), dtype=torch.float32, device=device - ).unsqueeze(-1) - - # 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, @@ -220,141 +218,65 @@ 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, - ) - - # 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"), - } - 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.") - - 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 = images[0] - render_overlay = overlay_mask( - image, - (render * 255.0).astype(np.uint8), - scale=1.0, + # 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 + }, ) - # difference left / right render / mask - difference = ( - cv2.cvtColor( - np.abs(render - masks[0].astype(np.float32) / 255.0), - cv2.COLOR_GRAY2BGR, + ) + 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, + on_iteration=on_iteration, + ) + rich.print("Entering optimization...") + result = diff_rendering_registration( + request=ImageRegistrationRequest( + cameras={ + "camera": CameraData( + intrinsics=intrinsics, ) - * 255.0 - ).astype(np.uint8) - # overlay segmentation mask - segmentation_overlay = overlay_mask( - image, - masks[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) - - # 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(images[i], (render * 255.0).astype(np.uint8), scale=1.0) - difference = np.abs(render - masks[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), + }, + robot_data=robot_data, + observations=observations, + 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( - 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_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 bb158ed..21ae4ce 100644 --- a/cli/rr_stereo_dr.py +++ b/cli/rr_stereo_dr.py @@ -1,40 +1,39 @@ 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, - TorchKinematics, - TorchMeshContainer, - VirtualCamera, -) from roboreg.io import ( find_files, load_robot_data_from_ros_xacro, load_robot_data_from_urdf_file, - parse_stereo_data, + parse_camera_info, + parse_stereo_observations, +) +from roboreg.registration.image.callbacks import RenderOverlayCallback +from roboreg.registration.image.config import ( + CameraConfig, + ConvergenceConfig, + DiffRenderingRegistrationConfig, + PlateauSchedulerConfig, +) +from roboreg.registration.image.objectives import ( + RenderingObjectiveType, + create_rendering_objective, +) +from roboreg.registration.image.request import CameraData, ImageRegistrationRequest +from roboreg.registration.image.solver import ( + DiffRenderingRegistration, + OptimizationCallback, + OptimizationState, ) -from roboreg.losses import soft_dice_loss -from roboreg.util import mask_distance_transform, mask_exponential_decay, overlay_mask 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 @@ -42,39 +41,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( + "--convergence-tolerance", + type=float, + default=ConvergenceConfig().tolerance, ) parser.add_argument( - "--step-size", + "--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", @@ -196,63 +208,36 @@ 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" 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_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=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), ) - # pre-process data - joint_states = torch.tensor( - np.array(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] - 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] - 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) - - # 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, @@ -268,207 +253,73 @@ 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, - ) - - # 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"), - } - 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() + # 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 + }, ) - else: - raise ValueError("Invalid registration mode.") - 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 = 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 = 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 - left_masks[0].astype(np.float32) / 255.0), - cv2.COLOR_GRAY2BGR, - ) - * 255.0 - ).astype(np.uint8) - ) - differences.append( - ( - cv2.cvtColor( - np.abs( - right_render - right_masks[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, - left_masks[0], - mode="b", - scale=1.0, - ) - ) - segmentation_overlays.append( - overlay_mask( - right_image, - right_masks[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, + 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, + on_iteration=[print_optimization_state], + ) + rich.print("Entering optimization...") + result = diff_rendering_registration( + request=ImageRegistrationRequest( + cameras={ + "left": CameraData( + intrinsics=left_intrinsics, ), - ) - 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( - 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 - ) - 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 - ) - - 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), + "right": CameraData( + intrinsics=right_intrinsics, + reference_to_camera=right_extrinsics, + ), + }, + robot_data=robot_data, + observations=observations, + 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( - 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/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 cc95da1..b524e5e 100644 --- a/roboreg/core/structs.py +++ b/roboreg/core/structs.py @@ -1,12 +1,19 @@ 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 class TorchMeshContainer: @@ -304,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 @@ -344,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 @@ -380,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/hydra_icp.py b/roboreg/hydra_icp.py deleted file mode 100644 index 616d6e9..0000000 --- a/roboreg/hydra_icp.py +++ /dev/null @@ -1,309 +0,0 @@ -from typing import List, Tuple - -import torch -from rich import print -from rich.progress import track - - -def kabsch_register( - input: torch.Tensor, target: 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. - - Args: - input (torch.Tensor): input of shape (..., M, 3). - target(torch.Tensor): target of shape (..., M, 3). - - Returns: - Tuple[torch.Tensor,torch.Tensor]: - - Rotation matrix of shape (..., 3, 3). - - Translation vector of shape (..., 3). - """ - # compute centroids - input_centroid = torch.mean(input, dim=-2) - target_centroid = torch.mean(target, dim=-2) - - # compute centered points - input_centered = input - input_centroid - target_centered = target - target_centroid - - # compute covariance matrix - H = target_centered.transpose(-1, -2) @ input_centered - - # compute SVD - U, _, V = torch.svd(H) - - E = torch.eye(3, dtype=U.dtype, device=U.device) - E[-1, -1] = torch.det(V @ U.transpose(-1, -2)) - - # compute rotation - R = V @ E @ U.transpose(-1, -2) - - # compute translation - t = target_centroid - input_centroid @ R - return R, t - - -def hydra_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. - - 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. - - 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: - raise ValueError("Input and target must have shape (..., 3).") - if max_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 - - -def hydra_centroid_alignment( - Xs: List[torch.Tensor], - Ys: List[torch.Tensor], -) -> torch.Tensor: - r"""Aligns centroids of Xs 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). - - Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ Xs = Ys. - """ - # 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] - - # estimate transform - R, t = kabsch_register( - torch.stack(Xs_centroids).unsqueeze(0), - torch.stack(Ys_centroids).unsqueeze(0), - ) - - HT = torch.eye(4, dtype=R.dtype, device=R.device) - R = R.squeeze(0) - t = t.squeeze(0) - HT[:3, :3] = R.T - HT[:3, 3] = t - return HT - - -def hydra_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: - 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. - - Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ observations = meshes. - """ - 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)): - # search correspondences - observations_tf = observations[i] @ HT[:3, :3].T + HT[:3, 3] - matchindices, mask = hydra_correspondence_indices( - observations_tf, meshes[i], max_distance - ) - - observation_corr.append(observations[i][mask]) - mesh_corr.append(meshes[i][matchindices[mask]].squeeze()) - - observation_corr = torch.concatenate(observation_corr).unsqueeze(0) - mesh_corr = torch.concatenate(mesh_corr).unsqueeze(0) - - ( - R, - t, - ) = kabsch_register( - observation_corr, - mesh_corr, - ) - R = R.squeeze(0) - t = t.squeeze(0) - HT[:3, :3] = R.T - HT[:3, 3] = t - - # compute rmse between observation and mesh_corr - rmse = torch.sqrt( - torch.mean( - torch.sum( - torch.pow( - mesh_corr - observation_corr, - 2, - ), - dim=-1, - ) - ) - ) - - if abs(prev_rmse - rmse.item()) < rmse_change and exit_early: - print("Converged early. Exiting.") - break - - prev_rmse = rmse.item() - - print("HT estimate:\n", HT) - return HT - - -def hydra_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: - 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. - - Returns: - torch.Tensor: Homogeneous transformation of shape (4, 4). HT @ observations = meshes. - """ - HT = HT_init # HT @ observation = mesh - - observations_cross_mat = [] - for i in range(len(observations)): - # build observation cross product matrix, refer eq. 4 (gets created once) - observations_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]), - ], - dim=-1, - ).reshape(-1, 3, 3) - ) - - # implementation of algorithm 1 - prev_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): - 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( - observations_tf, meshes[i], max_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()]) - - 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) - - for _ in range(inner_max_iter): - # ||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 - 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]), - ) - # weight associated with Huber loss - kappa = ( - 1.345 * torch.median(torch.abs(B - torch.median(B))) / 0.6745 - ) # eq. 26 - W = torch.where( - torch.abs(B) < kappa, - torch.ones_like(B), - torch.full_like(B, kappa) / torch.abs(B), - ) - - dTh_vec, resid, rank, singvals = torch.linalg.lstsq(W[:, None] * A, W * B) - dTh[0, 1] = -dTh_vec[2] - dTh[0, 2] = dTh_vec[1] - dTh[1, 0] = dTh_vec[2] - dTh[1, 2] = -dTh_vec[0] - dTh[2, 0] = -dTh_vec[1] - dTh[2, 1] = dTh_vec[0] - - dTh[0, 3] = dTh_vec[3] - dTh[1, 3] = dTh_vec[4] - dTh[2, 3] = dTh_vec[5] - - HT = HT @ torch.linalg.matrix_exp(dTh) - - # compute rmse between observation and mesh_corr - rmse = torch.sqrt( - torch.mean( - torch.sum( - torch.pow( - meshes_corr - observations_corr, - 2, - ), - dim=-1, - ) - ) - ) - - if abs(prev_rmse - rmse.item()) < rmse_change: - print("Converged early. Exiting.") - break - - prev_rmse = rmse.item() - - print("HT estimate:\n", HT) - return HT 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..c66576f 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.registration.image.request import CameraObservations, ImageObservations +from roboreg.registration.point_cloud.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,139 +348,175 @@ 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( - 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. +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. - mask_files (List[Path]): Mask files. + if image is None: + raise ValueError(f"Failed to read image '{path}'.") - Returns: - Tuple[List[np.ndarray],List[np.ndarray],List[np.ndarray]]: - - Images of shape HxWx3. - - Joint states. - - Masks of shape HxW. - """ - if len(image_files) != len(joint_states_files) or len(image_files) != len( - mask_files - ): - raise ValueError("Number of images, joint states, masks do not match.") + 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 + + +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], +) -> 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"Masks: {[f.name for f in mask_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 = ( + [_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", + ) - 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 - - -def parse_stereo_data( - 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. + return ImageObservations( + joint_states=joint_states, + cameras={ + "camera": CameraObservations( + images=images, + targets=targets, + ) + }, + ) - 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_mask_files (List[Path]): Left mask files. - 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. - """ - 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) - ): +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 masks: {[f.name for f in left_mask_files]}") - rich.print(f"Right masks: {[f.name for f in right_mask_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 + ) + + 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", + ) - 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] - 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 + 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, + cameras={ + "left": CameraObservations( + images=left_images, + targets=left_targets, + ), + "right": CameraObservations( + images=right_images, + targets=right_targets, + ), + }, + ) diff --git a/roboreg/io/robot_data.py b/roboreg/io/robot_data.py index 6e49c84..39e20b5 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.robot 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/registration/__init__.py b/roboreg/registration/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/roboreg/registration/_validation.py b/roboreg/registration/_validation.py new file mode 100644 index 0000000..e66ce3d --- /dev/null +++ b/roboreg/registration/_validation.py @@ -0,0 +1,41 @@ +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(f"Extrinsics must have shape (4, 4), 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_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/registration/image/__init__.py b/roboreg/registration/image/__init__.py new file mode 100644 index 0000000..e69de29 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 new file mode 100644 index 0000000..c63c494 --- /dev/null +++ b/roboreg/registration/image/config.py @@ -0,0 +1,70 @@ +import math +from dataclasses import dataclass, field +from typing import Literal, Tuple + + +@dataclass(frozen=True) +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.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 ConvergenceConfig: + max_iterations: int = 400 + tolerance: float = 1.0e-3 + patience: int = 100 + + 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: Literal["min", "max"] = "min" + factor: float = 0.1 + patience: int = 40 + 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 DiffRenderingRegistrationConfig: + camera: CameraConfig = field(default_factory=CameraConfig) + + optimizer: str = "AdamW" + lr: float = 3.0e-2 + + convergence: ConvergenceConfig = field(default_factory=ConvergenceConfig) + plateau_scheduler: PlateauSchedulerConfig = field( + default_factory=PlateauSchedulerConfig + ) + + def __post_init__(self) -> None: + if self.lr <= 0: + raise ValueError("lr must be positive.") diff --git a/roboreg/registration/image/objectives.py b/roboreg/registration/image/objectives.py new file mode 100644 index 0000000..86ec570 --- /dev/null +++ b/roboreg/registration/image/objectives.py @@ -0,0 +1,186 @@ +from dataclasses import dataclass +from enum import Enum +from typing import Protocol + +import numpy as np +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 + + 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. + 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 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.detach().cpu().numpy() + distance_maps = [mask_distance_transform(mask) for mask in targets_np] + return torch.as_tensor( + np.stack(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 + + 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. + 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 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.detach().cpu().numpy() + decay_maps = [ + mask_exponential_decay(mask, sigma=self._config.sigma) + for mask in targets_np + ] + return torch.as_tensor( + np.stack(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 + + 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.""" + + def __init__( + self, + config: ProbabilityMapConfig | None = None, + ) -> None: + self._config = config or ProbabilityMapConfig() + + 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.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() + + +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 new file mode 100644 index 0000000..facaf34 --- /dev/null +++ b/roboreg/registration/image/request.py @@ -0,0 +1,81 @@ +from dataclasses import dataclass + +import numpy as np + +from roboreg.core.robot import RobotData +from roboreg.registration._validation import ( + validate_extrinsics, + validate_images, + validate_intrinsics, + validate_targets, +) + + +@dataclass(frozen=True) +class CameraData: + intrinsics: np.ndarray + reference_to_camera: np.ndarray | None = None + + def __post_init__(self) -> None: + validate_intrinsics(self.intrinsics) + if self.reference_to_camera is not None: + validate_extrinsics(self.reference_to_camera) + + +@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( + "Expected the same number of images and targets, " + f"got {len(self.images)} and {len(self.targets)}." + ) + + 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.") + + if self.images is not None: + validate_images(self.images, "images") + + 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 image_shape != target_shape: + raise ValueError( + f"Image shape {image_shape} does not match " + f"target shape {target_shape}." + ) + + @property + def shape(self) -> tuple[int, int]: + return self.targets[0].shape[:2] + + +@dataclass(frozen=True) +class ImageObservations: + joint_states: list[np.ndarray] + cameras: dict[str, CameraObservations] + + +@dataclass(frozen=True) +class ImageRegistrationRequest: + cameras: dict[str, CameraData] + robot_data: RobotData + observations: ImageObservations + initial_extrinsics: np.ndarray + + def __post_init__(self) -> None: + 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 new file mode 100644 index 0000000..4a36ee8 --- /dev/null +++ b/roboreg/registration/image/solver.py @@ -0,0 +1,254 @@ +from dataclasses import dataclass +from typing import Callable, Iterable + +import numpy as np +import pytorch_kinematics as pk +import torch +import torch.nn.functional as F + +from roboreg.core.rendering import NVDiffRastRenderer +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 ( + ImageObservations, + ImageRegistrationRequest, +) +from roboreg.registration.result import RegistrationResult, TerminationReason +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) + self._on_iteration = on_iteration or [] + + def __call__( + self, + request: ImageRegistrationRequest, + ) -> RegistrationResult: + joint_states, preprocessed_targets = self._prepare_image_observations( + request.observations + ) + 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) + 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 + ) + # 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 + ) + return joint_states, preprocessed_targets + + 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: in case of + # hardware resolution and rendering resolution mismatch + 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, + current_resolution=native_resolution, + target_resolution=target_resolution, + ) + # prepare virtual camera + cameras[camera_name] = VirtualCamera( + resolution=target_resolution, + intrinsics=intrinsics, + extrinsics=camera_data.reference_to_camera, + 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), + ) + + 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: + 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") + 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: 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_losses[camera_name] = self._objective( + preprocessed_targets=preprocessed_targets[camera_name], + renders=render, + ) + 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) + 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, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) diff --git a/roboreg/registration/point_cloud/__init__.py b/roboreg/registration/point_cloud/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/roboreg/registration/point_cloud/config.py b/roboreg/registration/point_cloud/config.py new file mode 100644 index 0000000..79122a0 --- /dev/null +++ b/roboreg/registration/point_cloud/config.py @@ -0,0 +1,37 @@ +from dataclasses import dataclass, field + + +@dataclass(frozen=True) +class DepthToPointCloudConfig: + 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 + + depth_to_point_cloud: DepthToPointCloudConfig = field( + default_factory=DepthToPointCloudConfig + ) + + max_correspondence_distance: float = 0.1 + rmse_change_tolerance: float = 1e-6 + + +@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) + 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 new file mode 100644 index 0000000..c2a4d71 --- /dev/null +++ b/roboreg/registration/point_cloud/hydra.py @@ -0,0 +1,342 @@ +from typing import List, Tuple + +import torch +from rich.progress import track + +from roboreg.registration.result import RegistrationResult, TerminationReason + + +def kabsch_register( + 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 observed_vertices @ R + t = reference_vertices. + + Args: + 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]: + - Rotation matrix of shape (..., 3, 3). + - Translation vector of shape (..., 3). + """ + # compute centroids + observed_centroid = torch.mean(observed_vertices, dim=-2) + reference_centroid = torch.mean(reference_vertices, dim=-2) + + # compute centered points + observed_centered = observed_vertices - observed_centroid + reference_centered = reference_vertices - reference_centroid + + # compute covariance matrix + H = reference_centered.transpose(-1, -2) @ observed_centered + + # compute SVD + U, _, V = torch.svd(H) + + E = torch.eye(3, dtype=U.dtype, device=U.device) + E[-1, -1] = torch.det(V @ U.transpose(-1, -2)) + + # compute rotation + R = V @ E @ U.transpose(-1, -2) + + # compute translation + t = reference_centroid - observed_centroid @ R + return R, t + + +def correspondence_indices( + 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: + 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 observed_vertices.shape[-1] != 3 or reference_vertices.shape[-1] != 3: + raise ValueError("Input and target must have shape (..., 3).") + if max_correspondence_distance < 0: + raise ValueError("Max distance must be positive.") + 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( + observed_vertices: List[torch.Tensor], + reference_vertices: List[torch.Tensor], +) -> torch.Tensor: + r"""Aligns centroids of observed_vertices and Ys as an initial guess. + + Args: + 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 @ observed_vertices = reference_vertices. + """ + # for each cloud compute centroid + 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( + 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) + R = R.squeeze(0) + t = t.squeeze(0) + HT[:3, :3] = R.T + HT[:3, 3] = t + return HT + + +def point_to_point_icp( + HT_init: 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 @ 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: + RegistrationResult: Result with homogeneous transformation of shape (4, 4). HT @ observed_vertices = reference_vertices. + """ + HT = HT_init + # registration + 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 = observed_vertices[i] @ HT[:3, :3].T + HT[:3, 3] + match_indices, mask = correspondence_indices( + observations_tf, reference_vertices[i], max_correspondence_distance + ) + + observed_correspondences.append(observed_vertices[i][mask]) + reference_correspondences.append( + reference_vertices[i][match_indices[mask]].squeeze() + ) + + observed_correspondences = torch.concatenate( + observed_correspondences + ).unsqueeze(0) + reference_correspondences = torch.concatenate( + reference_correspondences + ).unsqueeze(0) + + ( + R, + t, + ) = kabsch_register( + observed_correspondences, + reference_correspondences, + ) + R = R.squeeze(0) + t = t.squeeze(0) + HT[:3, :3] = R.T + HT[:3, 3] = t + + # compute rmse between observed_correspondences and reference_correspondences + rmse = torch.sqrt( + torch.mean( + torch.sum( + torch.pow( + reference_correspondences - observed_correspondences, + 2, + ), + dim=-1, + ) + ) + ) + + if abs(previous_rmse - rmse.item()) < rmse_change_tolerance: + return RegistrationResult( + extrinsics=HT, + iterations=iteration, + termination_reason=TerminationReason.CONVERGED, + ) + + previous_rmse = rmse.item() + + return RegistrationResult( + extrinsics=HT, + iterations=max_iterations, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) + + +def point_to_plane_robust_icp( + HT_init: 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 @ 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: + RegistrationResult: Result with homogeneous transformation of shape (4, 4). HT @ observed_vertices = reference_vertices. + """ + HT = HT_init # HT @ observed_vertices = reference_vertices + + observed_cross_mat = [] + for i in range(len(observed_vertices)): + # build observation cross product matrix, refer eq. 4 (gets created once) + observed_cross_mat.append( + torch.stack( + [ + 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 + previous_rmse = float("inf") + dTh = torch.zeros_like(HT) + 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 + 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 + ) + + 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()] + ) + + 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(max_inner_iterations): + # ||A @ dTh - B||^2, refer eq. 14 + 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( + reference_normals_correspondences, + reference_correspondences + - (observed_correspondences @ HT[:3, :3].T + HT[:3, 3]), + ) + # weight associated with Huber loss + kappa = ( + 1.345 * torch.median(torch.abs(B - torch.median(B))) / 0.6745 + ) # eq. 26 + W = torch.where( + torch.abs(B) < kappa, + torch.ones_like(B), + torch.full_like(B, kappa) / torch.abs(B), + ) + + dTh_vec, resid, rank, singvals = torch.linalg.lstsq(W[:, None] * A, W * B) + dTh[0, 1] = -dTh_vec[2] + dTh[0, 2] = dTh_vec[1] + dTh[1, 0] = dTh_vec[2] + dTh[1, 2] = -dTh_vec[0] + dTh[2, 0] = -dTh_vec[1] + dTh[2, 1] = dTh_vec[0] + + dTh[0, 3] = dTh_vec[3] + dTh[1, 3] = dTh_vec[4] + dTh[2, 3] = dTh_vec[5] + + HT = HT @ torch.linalg.matrix_exp(dTh) + + # compute rmse between observation and mesh_correspondences + rmse = torch.sqrt( + torch.mean( + torch.sum( + torch.pow( + reference_correspondences - observed_correspondences, + 2, + ), + dim=-1, + ) + ) + ) + + if abs(previous_rmse - rmse.item()) < rmse_change_tolerance: + return RegistrationResult( + extrinsics=HT, + iterations=outer_iteration, + termination_reason=TerminationReason.CONVERGED, + ) + + previous_rmse = rmse.item() + + return RegistrationResult( + extrinsics=HT, + iterations=max_outer_iterations, + termination_reason=TerminationReason.MAX_ITERATIONS, + ) diff --git a/roboreg/registration/point_cloud/request.py b/roboreg/registration/point_cloud/request.py new file mode 100644 index 0000000..09966e7 --- /dev/null +++ b/roboreg/registration/point_cloud/request.py @@ -0,0 +1,57 @@ +from dataclasses import dataclass +from typing import List, Tuple + +import numpy as np + +from roboreg.core.robot import RobotData +from roboreg.registration._validation import ( + validate_intrinsics, + validate_masks, + validate_targets, +) + + +@dataclass(frozen=True) +class HydraObservations: + joint_states: List[np.ndarray] + masks: List[np.ndarray] + depths: List[np.ndarray] + + def __post_init__(self) -> None: + 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_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) +class HydraRequest: + intrinsics: np.ndarray + robot_data: RobotData + observations: HydraObservations + + def __post_init__(self) -> None: + validate_intrinsics(self.intrinsics) diff --git a/roboreg/registration/point_cloud/solver.py b/roboreg/registration/point_cloud/solver.py new file mode 100644 index 0000000..682f0bc --- /dev/null +++ b/roboreg/registration/point_cloud/solver.py @@ -0,0 +1,217 @@ +from dataclasses import dataclass +from typing import Callable, 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 + + +HydraCallback = Callable[ + [HydraProblem, RegistrationResult], + 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.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 + 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.depth_to_point_cloud.dilation_kernel_size, + config.depth_to_point_cloud.dilation_kernel_size, + ] + ), + erosion_kernel=np.ones( + [ + config.depth_to_point_cloud.erosion_kernel_size, + config.depth_to_point_cloud.erosion_kernel_size, + ] + ), + ) + if config.depth_to_point_cloud.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", + on_after_registration: HydraCallback | None = None, + ) -> None: + self._config = config or HydraICPConfig() + self._device = torch.device(device) + self._on_after_registration = on_after_registration + + 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 + ) + result = point_to_point_icp( + HT_init, + hydra_problem.observed_vertices, + hydra_problem.reference_vertices, + 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._on_after_registration is not None: + self._on_after_registration(hydra_problem, result) + return result + + +class HydraRobustICP: + def __init__( + self, + config: HydraRobustICPConfig | None = None, + device: torch.device | str = "cuda", + on_after_registration: HydraCallback | None = None, + ) -> None: + self._config = config or HydraRobustICPConfig() + self._device = torch.device(device) + self._on_after_registration = on_after_registration + + 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 + ) + result = point_to_plane_robust_icp( + HT_init, + hydra_problem.observed_vertices, + hydra_problem.reference_vertices, + hydra_problem.reference_normals, + 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._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 new file mode 100644 index 0000000..c7cdbb4 --- /dev/null +++ b/roboreg/registration/result.py @@ -0,0 +1,25 @@ +from dataclasses import dataclass +from enum import Enum + +import torch + + +class TerminationReason(str, Enum): + CONVERGED = "converged" + MAX_ITERATIONS = "max_iterations" + FAILED = "failed" + + def __str__(self) -> str: + return self.value + + +@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/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/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/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( 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 b3a13ea..09ff687 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,124 @@ 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"), + target_files=find_files(path, "mask_sam2_left_*.png"), ) assert ( - len(images) == len(joint_states) == len(masks) + len(observations.cameras["camera"].images) + == len(observations.joint_states) + == len(observations.cameras["camera"].targets) ), "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.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.cameras["camera"].targets[0] >= 0) and np.all( + observations.cameras["camera"].targets[0] <= 255 ), "Expected mask in range [0, 255]." assert ( - masks[0].shape[:2] == images[0].shape[:2] + observations.cameras["camera"].targets[0].shape[:2] + == observations.cameras["camera"].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_target_files=find_files(path, "mask_sam2_left_*.png"), + right_target_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.cameras["left"].images) + == len(observations.cameras["right"].images) + == len(observations.joint_states) + == 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(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 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.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.cameras["left"].targets[0].ndim == 2, "Expected 2D left mask." + assert ( + observations.cameras["left"].targets[0].dtype == np.uint8 + ), "Expected unsigned integers for left mask." + 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 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 ( - right_masks[0].dtype == np.uint8 + 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.cameras["right"].targets[0].ndim == 2, "Expected 2D right mask." + assert ( + observations.cameras["right"].targets[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.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 ( - left_masks[0].shape[:2] == 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 ( - right_masks[0].shape[:2] == 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." @@ -200,6 +227,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..42a9b79 100644 --- a/test/test_hydra_icp.py +++ b/test/test_hydra_icp.py @@ -6,19 +6,19 @@ 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, find_files, load_meshes, parse_camera_info, - apply_mesh_origins, - parse_hydra_data, + 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, @@ -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) @@ -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 - matchindices, mask = hydra_correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 # remove some elements randomly + observed_vertices = torch.rand(M, dim) + reference_vertices = torch.rand(N, dim) # e.g. the mesh vertices + matchindices, mask = correspondence_indices( + 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) - matchindices, mask = hydra_correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 + observed_vertices = torch.rand(batch_size, M, dim) + reference_vertices = torch.rand(batch_size, N, dim) + matchindices, mask = correspondence_indices( + observed_vertices, + reference_vertices, + max_correspondence_distance=np.sqrt(dim) / 2.0, ) test_index_shape(matchindices, mask, torch.Size([batch_size, M]), N) @@ -96,16 +100,18 @@ def test_index_shape( M = 10 N = 100 - input = torch.rand(M, dim) - target = torch.rand(N, dim) - matchindices, mask = hydra_correspondence_indices( - input, target, max_distance=np.sqrt(dim) / 2.0 + observed_vertices = torch.rand(M, dim) + reference_vertices = torch.rand(N, dim) + matchindices, mask = correspondence_indices( + observed_vertices, + reference_vertices, + max_correspondence_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" @@ -118,7 +124,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 +145,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( @@ -156,19 +162,19 @@ def test_hydra_icp(): ) # perform forward kinematics - mesh_vertices = meshes.vertices.clone() + reference_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[ + 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,11 +182,13 @@ def test_hydra_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) - 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 @@ -194,8 +202,8 @@ def test_hydra_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 = [ @@ -204,39 +212,41 @@ 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 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 = hydra_centroid_alignment(observed_vertices, mesh_vertices) - HT = hydra_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.") -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" @@ -249,7 +259,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 +280,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( @@ -287,19 +297,19 @@ def test_hydra_robust_icp() -> None: ) # perform forward kinematics - mesh_vertices = meshes.vertices.clone() + reference_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(): - 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 @@ -310,7 +320,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 @@ -325,12 +337,12 @@ def test_hydra_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 @@ -340,38 +352,40 @@ 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 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 = hydra_centroid_alignment(observed_vertices, mesh_vertices) - HT = hydra_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__": @@ -383,5 +397,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() diff --git a/test/util/test_mask.py b/test/util/test_mask.py index 651573c..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,124 +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, - ) - dilated_mask = mask_dilate_with_kernel(mask) - cv2.imshow("mask", mask) - cv2.imshow("dilated_mask", dilated_mask) - cv2.waitKey(0) - cv2.destroyAllWindows() - - -@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, +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, ) + kernel = np.ones((3, 3), dtype=np.uint8) + + result = mask_dilate_with_kernel(mask, kernel) - # 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.imshow("mask", mask) - cv2.imshow("distance_map", distance_map) - cv2.waitKey(0) - cv2.destroyAllWindows() - - # 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.imshow("inverse_mask", inverse_mask) - cv2.imshow("inverse_distance_map", inverse_distance_map) - cv2.waitKey(0) - cv2.destroyAllWindows() - - -@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, + expected = np.zeros((7, 7), dtype=np.uint8) + expected[1:6, 1:6] = 255 + + np.testing.assert_array_equal(result, expected) + assert result.dtype == np.uint8 + + +@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.imshow("mask", mask) - cv2.imshow("eroded_mask", eroded_mask) - cv2.waitKey(0) - cv2.destroyAllWindows() + 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.imshow("mask", mask) - cv2.imshow("exponential_decay", exponential_decay) - cv2.waitKey(0) - cv2.destroyAllWindows() + 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.imshow("mask", mask) - cv2.imshow("boundary_mask", boundary_mask) - cv2.imshow("overlay", overlay) - cv2.waitKey(0) - cv2.destroyAllWindows() + + 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]) + kernel = np.ones((3, 3), dtype=np.uint8) + + result = mask_extract_extended_boundary( + mask, + dilation_kernel=kernel, + erosion_kernel=kernel, ) - 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() - - -if __name__ == "__main__": - test_dilate_with_kernel() - test_distance_transform() - test_erode_with_kernel() - test_exponential_decay() - test_extract_boundary() - test_extract_extended_boundary() + + 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) + + with pytest.raises(ValueError, match="Expected a 2D mask"): + function(mask) 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()