use core::f64;
use kiddo::immutable::float::kdtree::ImmutableKdTree;
use super::ops::{find_correspondences, fit_transformation, update_transformation};
use crate::{linalg::transform_points3d, pointcloud::PointCloud};
#[derive(Debug, Clone)]
pub struct ICPResult {
pub rotation: [[f64; 3]; 3],
pub translation: [f64; 3],
pub num_iterations: usize,
pub rmse: f64,
}
#[derive(Debug, Clone)]
pub struct ICPConvergenceCriteria {
pub max_iterations: usize,
pub tolerance: f64,
}
pub fn icp_vanilla(
source: &PointCloud,
target: &PointCloud,
initial_rot: [[f64; 3]; 3],
initial_trans: [f64; 3],
criteria: ICPConvergenceCriteria,
) -> Result<ICPResult, Box<dyn std::error::Error>> {
let mut result = ICPResult {
rotation: initial_rot,
translation: initial_trans,
num_iterations: 0,
rmse: f64::INFINITY,
};
let kdtree: ImmutableKdTree<f64, u32, 3, 32> = ImmutableKdTree::new_from_slice(target.points());
let mut transformed_points = vec![[0.0; 3]; source.points().len()];
transform_points3d(
source.points(),
&result.rotation,
&result.translation,
&mut transformed_points,
)?;
let mut current_source = transformed_points;
for i in 0..criteria.max_iterations {
log::debug!("Iteration: {i}");
let now = std::time::Instant::now();
let (current_source_match, current_target_match, distances) =
find_correspondences(¤t_source, target.points(), &kdtree);
log::debug!(
"Num correspondences: {}-{}",
current_source_match.len(),
current_target_match.len()
);
let mut rr_delta = [[0.0; 3]; 3];
let mut tt_delta = [0.0; 3];
fit_transformation(
¤t_source_match,
¤t_target_match,
&mut rr_delta,
&mut tt_delta,
);
let mut transformed_points = vec![[0.0; 3]; current_source.len()];
transform_points3d(
¤t_source,
&rr_delta,
&tt_delta,
&mut transformed_points,
)?;
update_transformation(
&mut result.rotation,
&mut result.translation,
&rr_delta,
&tt_delta,
);
let rmse = (distances.iter().sum::<f64>() / distances.len() as f64).sqrt();
result.num_iterations += 1;
if (result.rmse - rmse).abs() < criteria.tolerance {
log::debug!("ICP converged in {i} iterations with error {rmse}");
result.rmse = rmse;
break;
}
result.rmse = rmse;
current_source = transformed_points;
let elapsed = now.elapsed();
log::debug!("elapsed: {elapsed:?}");
}
Ok(result)
}
#[cfg(test)]
mod tests {
use super::{icp_vanilla, ICPConvergenceCriteria};
use crate::{
linalg::transform_points3d, pointcloud::PointCloud,
transforms::axis_angle_to_rotation_matrix,
};
#[test]
fn test_icp_vanilla() -> Result<(), Box<dyn std::error::Error>> {
let num_points = 100;
let points_src = (0..num_points)
.map(|_| {
[
rand::random::<f64>(),
rand::random::<f64>(),
rand::random::<f64>(),
]
})
.collect::<Vec<_>>();
let dst_r_src = axis_angle_to_rotation_matrix(&[1.0, 0.0, 0.0], 0.1)?;
let dst_t_src = [0.1, 0.1, 0.1];
let mut points_dst = vec![[0.0; 3]; points_src.len()];
transform_points3d(&points_src, &dst_r_src, &dst_t_src, &mut points_dst)?;
let src_pcl = PointCloud::new(points_src, None, None);
let dst_pcl = PointCloud::new(points_dst, None, None);
let initial_rot = [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0]];
let initial_trans = [0.0, 0.0, 0.0];
let result = icp_vanilla(
&src_pcl,
&dst_pcl,
initial_rot,
initial_trans,
ICPConvergenceCriteria {
max_iterations: 100,
tolerance: 1e-6,
},
)?;
println!("result: {result:?}");
let mut r_error = [[0.0; 3]; 3];
for (i, r_error_row) in r_error.iter_mut().enumerate() {
for (j, r_error_cell) in r_error_row.iter_mut().enumerate() {
*r_error_cell = result.rotation[0][i] * dst_r_src[0][j]
+ result.rotation[1][i] * dst_r_src[1][j]
+ result.rotation[2][i] * dst_r_src[2][j];
}
}
let trace = r_error[0][0] + r_error[1][1] + r_error[2][2];
let angular_error = ((trace - 1.0) / 2.0).clamp(-1.0, 1.0).acos();
let translation_error = ((result.translation[0] - dst_t_src[0]).powi(2)
+ (result.translation[1] - dst_t_src[1]).powi(2)
+ (result.translation[2] - dst_t_src[2]).powi(2))
.sqrt();
assert!(
angular_error < 0.1,
"Angular rotation error too large: {} rad (expected < 0.1 rad)",
angular_error
);
assert!(
translation_error < 0.1,
"Translation error too large: {} (expected < 0.1)",
translation_error
);
assert!(
result.rmse < 1e-8,
"ICP did not converge to low error: RMSE = {}",
result.rmse
);
Ok(())
}
}