use crate::pose::fundamental::sampson_distance;
use crate::ransac::{RobustKernel, RobustKernelKind};
use faer::prelude::SpSolver;
use kornia_algebra::{Mat3F64, Vec2F64, Vec3F64};
#[derive(Clone, Copy, Debug)]
pub struct LmPoseConfig {
pub max_iters: usize,
pub initial_lambda: f64,
pub lambda_up: f64,
pub lambda_down: f64,
pub gradient_tol: f64,
pub step_tol: f64,
pub cost_tol: f64,
pub robust: RobustKernelKind,
pub robust_scale_sq: f64,
}
impl Default for LmPoseConfig {
fn default() -> Self {
Self {
max_iters: 10,
initial_lambda: 1e-3,
lambda_up: 10.0,
lambda_down: 0.5,
gradient_tol: 1e-9,
step_tol: 1e-9,
cost_tol: 1e-12,
robust: RobustKernelKind::Identity,
robust_scale_sq: f64::INFINITY,
}
}
}
#[inline]
pub(crate) fn hat(v: Vec3F64) -> Mat3F64 {
Mat3F64::from_cols(
Vec3F64::new(0.0, v.z, -v.y),
Vec3F64::new(-v.z, 0.0, v.x),
Vec3F64::new(v.y, -v.x, 0.0),
)
}
#[inline]
fn so3_exp(w: Vec3F64) -> Mat3F64 {
let theta_sq = w.x * w.x + w.y * w.y + w.z * w.z;
let theta = theta_sq.sqrt();
let wx = hat(w);
if theta < 1e-8 {
return Mat3F64::IDENTITY + wx;
}
let a = theta.sin() / theta;
let b = (1.0 - theta.cos()) / theta_sq;
let wx_sq = wx * wx;
Mat3F64::IDENTITY + wx * a + wx_sq * b
}
#[inline]
fn cross(a: Vec3F64, b: Vec3F64) -> Vec3F64 {
Vec3F64::new(
a.y * b.z - a.z * b.y,
a.z * b.x - a.x * b.z,
a.x * b.y - a.y * b.x,
)
}
#[inline]
fn tangent_basis(t: Vec3F64) -> (Vec3F64, Vec3F64) {
let ax = t.x.abs();
let ay = t.y.abs();
let az = t.z.abs();
let seed = if ax <= ay && ax <= az {
Vec3F64::new(1.0, 0.0, 0.0)
} else if ay <= az {
Vec3F64::new(0.0, 1.0, 0.0)
} else {
Vec3F64::new(0.0, 0.0, 1.0)
};
let b1 = cross(t, seed).normalize();
let b2 = cross(t, b1); (b1, b2)
}
#[inline]
pub(crate) fn fundamental_from_rt(
r: &Mat3F64,
t: &Vec3F64,
k1_inv: &Mat3F64,
k2_inv_t: &Mat3F64,
) -> Mat3F64 {
let e = hat(*t) * *r;
*k2_inv_t * e * *k1_inv
}
#[inline]
fn sum_sampson(f: &Mat3F64, x1: &[Vec2F64], x2: &[Vec2F64]) -> f64 {
let mut s = 0.0;
for (p1, p2) in x1.iter().zip(x2.iter()) {
s += sampson_distance(f, p1, p2);
}
s
}
#[inline]
fn residuals(f: &Mat3F64, x1: &[Vec2F64], x2: &[Vec2F64], out: &mut [f64]) {
for (i, (p1, p2)) in x1.iter().zip(x2.iter()).enumerate() {
let d = sampson_distance(f, p1, p2);
out[i] = if d > 0.0 { d.sqrt() } else { 0.0 };
}
}
pub fn refine_pose_lm(
r: Mat3F64,
t: Vec3F64,
x1_inl: &[Vec2F64],
x2_inl: &[Vec2F64],
k1: &Mat3F64,
k2: &Mat3F64,
cfg: &LmPoseConfig,
) -> (Mat3F64, Vec3F64) {
let n = x1_inl.len();
if n < 6 || x2_inl.len() != n {
return (r, t);
}
let k1_inv = k1.inverse();
let k2_inv_t = k2.inverse().transpose();
let mut r_cur = r;
let mut t_cur = t.normalize();
let f_cur = fundamental_from_rt(&r_cur, &t_cur, &k1_inv, &k2_inv_t);
let mut cost_cur = sum_sampson(&f_cur, x1_inl, x2_inl);
let cost_initial = cost_cur;
let mut lambda = cfg.initial_lambda;
let mut res = vec![0.0_f64; n];
let mut res_pert = vec![0.0_f64; n];
let mut res_minus = vec![0.0_f64; n];
let mut jac = vec![0.0_f64; n * 5];
for _iter in 0..cfg.max_iters {
let f_at_cur = fundamental_from_rt(&r_cur, &t_cur, &k1_inv, &k2_inv_t);
residuals(&f_at_cur, x1_inl, x2_inl, &mut res);
let (b1, b2) = tangent_basis(t_cur);
let h_fd = 1e-6_f64;
let two_h_inv = 1.0 / (2.0 * h_fd);
for col in 0..5 {
let (r_p, t_p) = apply_tangent_step(r_cur, t_cur, b1, b2, col, h_fd);
let f_p = fundamental_from_rt(&r_p, &t_p, &k1_inv, &k2_inv_t);
residuals(&f_p, x1_inl, x2_inl, &mut res_pert);
let (r_m, t_m) = apply_tangent_step(r_cur, t_cur, b1, b2, col, -h_fd);
let f_m = fundamental_from_rt(&r_m, &t_m, &k1_inv, &k2_inv_t);
residuals(&f_m, x1_inl, x2_inl, &mut res_minus);
let col_base = col * n;
for i in 0..n {
jac[col_base + i] = (res_pert[i] - res_minus[i]) * two_h_inv;
}
}
let kernel = cfg.robust;
let c_sq = cfg.robust_scale_sq;
let mut weights = vec![0.0_f64; n];
for i in 0..n {
weights[i] = kernel.weight(res[i] * res[i], c_sq);
}
let mut h_mat = [[0.0_f64; 5]; 5];
let mut g_vec = [0.0_f64; 5];
for c1 in 0..5 {
let col1 = &jac[c1 * n..c1 * n + n];
let mut g = 0.0;
for i in 0..n {
g += col1[i] * res[i] * weights[i];
}
g_vec[c1] = g;
for c2 in c1..5 {
let col2 = &jac[c2 * n..c2 * n + n];
let mut h = 0.0;
for i in 0..n {
h += col1[i] * col2[i] * weights[i];
}
h_mat[c1][c2] = h;
h_mat[c2][c1] = h;
}
}
let grad_inf = g_vec.iter().fold(0.0_f64, |m, &v| m.max(v.abs()));
if grad_inf < cfg.gradient_tol {
break;
}
let mut a = faer::Mat::<f64>::zeros(5, 5);
let mut rhs = faer::Mat::<f64>::zeros(5, 1);
for r_i in 0..5 {
for c_i in 0..5 {
let mut v = h_mat[r_i][c_i];
if r_i == c_i {
v += lambda * h_mat[r_i][r_i].abs().max(1e-12);
}
unsafe {
a.write_unchecked(r_i, c_i, v);
}
}
unsafe {
rhs.write_unchecked(r_i, 0, -g_vec[r_i]);
}
}
let sol = a.partial_piv_lu().solve(&rhs);
let delta = [
sol.read(0, 0),
sol.read(1, 0),
sol.read(2, 0),
sol.read(3, 0),
sol.read(4, 0),
];
if !delta.iter().all(|x: &f64| x.is_finite()) {
lambda *= cfg.lambda_up;
continue;
}
let step_norm = (delta.iter().map(|x| x * x).sum::<f64>()).sqrt();
if step_norm < cfg.step_tol {
break;
}
let (r_try, t_try) = retract(r_cur, t_cur, b1, b2, delta);
let f_try = fundamental_from_rt(&r_try, &t_try, &k1_inv, &k2_inv_t);
let cost_try = sum_sampson(&f_try, x1_inl, x2_inl);
if cost_try.is_finite() && cost_try < cost_cur {
let rel_drop = (cost_cur - cost_try).abs() / cost_cur.abs().max(1e-300);
r_cur = r_try;
t_cur = t_try;
cost_cur = cost_try;
lambda *= cfg.lambda_down;
if rel_drop < cfg.cost_tol {
break;
}
} else {
lambda *= cfg.lambda_up;
if lambda > 1e16 {
break;
}
}
}
if cost_cur.is_finite() && cost_cur <= cost_initial {
(r_cur, t_cur)
} else {
(r, t.normalize())
}
}
#[inline]
fn apply_tangent_step(
r: Mat3F64,
t: Vec3F64,
b1: Vec3F64,
b2: Vec3F64,
col: usize,
h: f64,
) -> (Mat3F64, Vec3F64) {
let mut delta = [0.0_f64; 5];
delta[col] = h;
retract(r, t, b1, b2, delta)
}
#[inline]
fn retract(
r: Mat3F64,
t: Vec3F64,
b1: Vec3F64,
b2: Vec3F64,
delta: [f64; 5],
) -> (Mat3F64, Vec3F64) {
let omega = Vec3F64::new(delta[0], delta[1], delta[2]);
let r_new = r * so3_exp(omega);
let b1s = Vec3F64::new(b1.x * delta[3], b1.y * delta[3], b1.z * delta[3]);
let b2s = Vec3F64::new(b2.x * delta[4], b2.y * delta[4], b2.z * delta[4]);
let t_raw = Vec3F64::new(
t.x + b1s.x + b2s.x,
t.y + b1s.y + b2s.y,
t.z + b1s.z + b2s.z,
);
let t_new = t_raw.normalize();
(r_new, t_new)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_so3_exp_identity() {
let m = so3_exp(Vec3F64::ZERO);
let diff: [f64; 9] = (m - Mat3F64::IDENTITY).into();
for d in diff {
assert!(d.abs() < 1e-15);
}
}
#[test]
fn test_so3_exp_small_angle() {
let w = Vec3F64::new(1e-10, -2e-10, 5e-11);
let m = so3_exp(w);
let expected = Mat3F64::IDENTITY + hat(w);
let d: [f64; 9] = (m - expected).into();
for v in d {
assert!(v.abs() < 1e-18);
}
}
#[test]
fn test_tangent_basis_orthonormal() {
let t = Vec3F64::new(0.2422, -0.2330, 0.9418).normalize();
let (b1, b2) = tangent_basis(t);
assert!(t.dot(b1).abs() < 1e-12);
assert!(t.dot(b2).abs() < 1e-12);
assert!(b1.dot(b2).abs() < 1e-12);
assert!((b1.length() - 1.0).abs() < 1e-12);
assert!((b2.length() - 1.0).abs() < 1e-12);
}
}