Skip to main content

spatialrust_vision/
pnp.rs

1//! Absolute pose from 3D–2D correspondences (PnP) with deterministic RANSAC.
2
3use spatialrust_math::{
4    solve_linear_system, symmetric_eigen3, LeastSquaresResult, Mat3, Vec2, Vec3,
5};
6
7use crate::{
8    AbsolutePose, CameraMatrix3, GeometricEstimate, ObjectImageCorrespondence,
9    RobustEstimationOptions, VisionError, VisionResult,
10};
11
12/// Estimates an object-to-camera pose from at least four correspondences.
13///
14/// Uses a calibrated DLT initialization followed by Gauss–Newton refinement on
15/// the SE(3) tangent space. Four points are accepted for the final refine path
16/// when an initial pose is recoverable; RANSAC minimal samples use six points.
17pub fn solve_pnp(
18    correspondences: &[ObjectImageCorrespondence],
19    camera: CameraMatrix3,
20) -> VisionResult<AbsolutePose> {
21    if correspondences.len() < 4 {
22        return Err(VisionError::InvalidParameter(
23            "PnP requires at least four object-image correspondences".into(),
24        ));
25    }
26    let initial = estimate_pnp_dlt(correspondences, camera)?;
27    refine_pnp(correspondences, camera, initial)
28}
29
30/// Robust PnP using deterministic six-point RANSAC and inlier refinement.
31pub fn solve_pnp_ransac(
32    correspondences: &[ObjectImageCorrespondence],
33    camera: CameraMatrix3,
34    options: RobustEstimationOptions,
35) -> VisionResult<GeometricEstimate<AbsolutePose>> {
36    let options = options.validate()?;
37    const SAMPLE: usize = 6;
38    if correspondences.len() < SAMPLE {
39        return Err(VisionError::InvalidParameter(
40            "robust PnP requires at least six correspondences".into(),
41        ));
42    }
43    let mut rng = XorShift64::new(options.seed);
44    let mut best: Option<(AbsolutePose, Vec<bool>, Vec<f64>, usize, f64)> = None;
45    let mut iteration_limit = options.max_iterations;
46    let mut iteration = 0;
47    while iteration < iteration_limit {
48        let indices = sample_unique(&mut rng, correspondences.len(), SAMPLE);
49        let sample = indices.iter().map(|&index| correspondences[index]).collect::<Vec<_>>();
50        if let Ok(model) =
51            estimate_pnp_dlt(&sample, camera).and_then(|pose| refine_pnp(&sample, camera, pose))
52        {
53            let residuals = correspondences
54                .iter()
55                .copied()
56                .map(|pair| pnp_residual(model, pair, camera))
57                .collect::<Vec<_>>();
58            let inliers =
59                residuals.iter().map(|&value| value <= options.threshold).collect::<Vec<_>>();
60            let count = inliers.iter().filter(|&&value| value).count();
61            let error = residuals
62                .iter()
63                .zip(&inliers)
64                .filter_map(|(value, &is_inlier)| is_inlier.then_some(*value))
65                .sum::<f64>();
66            let improves = best.as_ref().map_or(true, |candidate| {
67                count > candidate.3 || (count == candidate.3 && error < candidate.4)
68            });
69            if improves {
70                if count >= SAMPLE {
71                    let inlier_ratio = count as f64 / correspondences.len() as f64;
72                    let success = inlier_ratio.powi(SAMPLE as i32).clamp(0.0, 1.0);
73                    if success > 0.0 && success < 1.0 {
74                        let required = ((1.0 - options.confidence).ln() / (1.0 - success).ln())
75                            .ceil()
76                            .max(1.0) as usize;
77                        iteration_limit = iteration_limit.min(required.max(iteration + 1));
78                    } else if success == 1.0 {
79                        iteration_limit = iteration + 1;
80                    }
81                }
82                best = Some((model, inliers, residuals, count, error));
83            }
84        }
85        iteration += 1;
86    }
87    let (_, best_inliers, _, count, _) = best
88        .ok_or_else(|| VisionError::InvalidParameter("robust PnP found no valid model".into()))?;
89    if count < 4 {
90        return Err(VisionError::InvalidParameter("robust PnP found too few inliers".into()));
91    }
92    let inlier_pairs = correspondences
93        .iter()
94        .zip(&best_inliers)
95        .filter_map(|(&pair, &is_inlier)| is_inlier.then_some(pair))
96        .collect::<Vec<_>>();
97    let refined = solve_pnp(&inlier_pairs, camera)?;
98    let residuals = correspondences
99        .iter()
100        .copied()
101        .map(|pair| pnp_residual(refined, pair, camera))
102        .collect::<Vec<_>>();
103    let inliers = residuals.iter().map(|&value| value <= options.threshold).collect();
104    GeometricEstimate::try_new(refined, correspondences.len(), inliers, residuals)
105}
106
107/// Projects an object point with an absolute pose into pixel coordinates.
108pub fn project_object_point(
109    pose: AbsolutePose,
110    camera: CameraMatrix3,
111    object: Vec3<f64>,
112) -> VisionResult<Vec2<f64>> {
113    let camera_point = pose.transform_point(object);
114    if camera_point.z <= 1e-12 {
115        return Err(VisionError::InvalidParameter(
116            "projected point lies behind or on the camera plane".into(),
117        ));
118    }
119    let normalized =
120        Vec3::new(camera_point.x / camera_point.z, camera_point.y / camera_point.z, 1.0);
121    let pixel = camera.matrix().mul_vec3(normalized);
122    Ok(Vec2 { x: pixel.x / pixel.z, y: pixel.y / pixel.z })
123}
124
125fn estimate_pnp_dlt(
126    correspondences: &[ObjectImageCorrespondence],
127    camera: CameraMatrix3,
128) -> VisionResult<AbsolutePose> {
129    if correspondences.len() < 4 {
130        return Err(VisionError::InvalidParameter(
131            "PnP DLT requires at least four correspondences".into(),
132        ));
133    }
134    let mut normal = vec![vec![0.0; 12]; 12];
135    for pair in correspondences {
136        let object = pair.object();
137        let image = pair.image();
138        let rows = [
139            [
140                object.x,
141                object.y,
142                object.z,
143                1.0,
144                0.0,
145                0.0,
146                0.0,
147                0.0,
148                -image.x * object.x,
149                -image.x * object.y,
150                -image.x * object.z,
151                -image.x,
152            ],
153            [
154                0.0,
155                0.0,
156                0.0,
157                0.0,
158                object.x,
159                object.y,
160                object.z,
161                1.0,
162                -image.y * object.x,
163                -image.y * object.y,
164                -image.y * object.z,
165                -image.y,
166            ],
167        ];
168        for row in rows {
169            for first in 0..12 {
170                for second in 0..12 {
171                    normal[first][second] += row[first] * row[second];
172                }
173            }
174        }
175    }
176    let vector = smallest_symmetric_eigenvector(normal).ok_or_else(|| {
177        VisionError::InvalidParameter("PnP correspondences are degenerate".into())
178    })?;
179    let projection = Mat3::from_rows(
180        [vector[0], vector[1], vector[2]],
181        [vector[4], vector[5], vector[6]],
182        [vector[8], vector[9], vector[10]],
183    );
184    let translation_part = Vec3::new(vector[3], vector[7], vector[11]);
185    let calibrated = camera.inverse().mul_mat3(projection);
186    let calibrated_t = camera.inverse().mul_vec3(translation_part);
187    let (rotation, scale) = orthonormalize_rotation(calibrated)?;
188    let translation =
189        Vec3::new(calibrated_t.x / scale, calibrated_t.y / scale, calibrated_t.z / scale);
190    // Flip if most points have negative depth.
191    let pose = AbsolutePose::try_new(rotation, translation)?;
192    let positive =
193        correspondences.iter().filter(|pair| pose.transform_point(pair.object()).z > 0.0).count();
194    if positive * 2 < correspondences.len() {
195        AbsolutePose::try_new(
196            Mat3::from_rows(
197                [-rotation.m[0][0], -rotation.m[0][1], -rotation.m[0][2]],
198                [-rotation.m[1][0], -rotation.m[1][1], -rotation.m[1][2]],
199                [-rotation.m[2][0], -rotation.m[2][1], -rotation.m[2][2]],
200            ),
201            Vec3::new(-translation.x, -translation.y, -translation.z),
202        )
203    } else {
204        Ok(pose)
205    }
206}
207
208fn refine_pnp(
209    correspondences: &[ObjectImageCorrespondence],
210    camera: CameraMatrix3,
211    mut pose: AbsolutePose,
212) -> VisionResult<AbsolutePose> {
213    for _ in 0..20 {
214        let mut normal = vec![vec![0.0; 6]; 6];
215        let mut rhs = vec![0.0; 6];
216        let mut residual_sum = 0.0;
217        for pair in correspondences {
218            let camera_point = pose.transform_point(pair.object());
219            if camera_point.z <= 1e-12 {
220                continue;
221            }
222            let predicted = project_object_point(pose, camera, pair.object())?;
223            let error = [predicted.x - pair.image().x, predicted.y - pair.image().y];
224            residual_sum += error[0] * error[0] + error[1] * error[1];
225            let jacobian = projection_jacobian(pose, camera, pair.object(), camera_point)?;
226            for row in 0..2 {
227                for col in 0..6 {
228                    rhs[col] -= jacobian[row][col] * error[row];
229                    for other in 0..6 {
230                        normal[col][other] += jacobian[row][col] * jacobian[row][other];
231                    }
232                }
233            }
234        }
235        let LeastSquaresResult::Solved(delta) = solve_linear_system(normal, rhs) else {
236            break;
237        };
238        let update = Vec3::new(delta[0], delta[1], delta[2]);
239        if update.length() + Vec3::new(delta[3], delta[4], delta[5]).length() < 1e-10 {
240            break;
241        }
242        let rotated = exp_so3(update).mul_mat3(pose.rotation());
243        let (rotation, _) = orthonormalize_rotation(rotated)?;
244        let translation = pose.translation() + Vec3::new(delta[3], delta[4], delta[5]);
245        pose = AbsolutePose::try_new(rotation, translation)?;
246        if residual_sum < 1e-18 {
247            break;
248        }
249    }
250    Ok(pose)
251}
252
253fn projection_jacobian(
254    pose: AbsolutePose,
255    camera: CameraMatrix3,
256    object: Vec3<f64>,
257    camera_point: Vec3<f64>,
258) -> VisionResult<[[f64; 6]; 2]> {
259    let z = camera_point.z;
260    let z2 = z * z;
261    let fx = camera.matrix().m[0][0];
262    let fy = camera.matrix().m[1][1];
263    // d(u,v)/dX_c
264    let du_dx = fx / z;
265    let du_dy = 0.0;
266    let du_dz = -fx * camera_point.x / z2;
267    let dv_dx = 0.0;
268    let dv_dy = fy / z;
269    let dv_dz = -fy * camera_point.y / z2;
270    // dX_c / d(omega,t): omega acts as [omega]_x R X, t is additive.
271    let rotated = pose.rotation().mul_vec3(object);
272    let dx_domega = [
273        Vec3::new(0.0, -rotated.z, rotated.y),
274        Vec3::new(rotated.z, 0.0, -rotated.x),
275        Vec3::new(-rotated.y, rotated.x, 0.0),
276    ];
277    let mut jacobian = [[0.0; 6]; 2];
278    for axis in 0..3 {
279        let d = dx_domega[axis];
280        jacobian[0][axis] = du_dx * d.x + du_dy * d.y + du_dz * d.z;
281        jacobian[1][axis] = dv_dx * d.x + dv_dy * d.y + dv_dz * d.z;
282    }
283    jacobian[0][3] = du_dx;
284    jacobian[0][4] = du_dy;
285    jacobian[0][5] = du_dz;
286    jacobian[1][3] = dv_dx;
287    jacobian[1][4] = dv_dy;
288    jacobian[1][5] = dv_dz;
289    Ok(jacobian)
290}
291
292fn pnp_residual(pose: AbsolutePose, pair: ObjectImageCorrespondence, camera: CameraMatrix3) -> f64 {
293    match project_object_point(pose, camera, pair.object()) {
294        Ok(pixel) => (pixel.x - pair.image().x).hypot(pixel.y - pair.image().y),
295        Err(_) => f64::MAX,
296    }
297}
298
299fn orthonormalize_rotation(matrix: Mat3<f64>) -> VisionResult<(Mat3<f64>, f64)> {
300    let eigen = symmetric_eigen3(matrix.transpose().mul_mat3(matrix));
301    let scale = ((eigen.eigenvalues[0].max(0.0).sqrt()
302        + eigen.eigenvalues[1].max(0.0).sqrt()
303        + eigen.eigenvalues[2].max(0.0).sqrt())
304        / 3.0)
305        .max(1e-12);
306    let right = eigen.eigenvectors;
307    let mut left_cols = [Vec3::new(0.0, 0.0, 0.0); 3];
308    for (column, left_col) in left_cols.iter_mut().enumerate() {
309        let right_col = Vec3::new(right.m[0][column], right.m[1][column], right.m[2][column]);
310        let sigma = eigen.eigenvalues[column].max(0.0).sqrt().max(1e-12);
311        *left_col = scale_vec(matrix.mul_vec3(right_col), 1.0 / (sigma * scale));
312    }
313    left_cols[0] = left_cols[0].normalize();
314    left_cols[1] =
315        (left_cols[1] - scale_vec(left_cols[0], left_cols[0].dot(left_cols[1]))).normalize();
316    left_cols[2] = left_cols[0].cross(left_cols[1]).normalize();
317    let right0 = Vec3::new(right.m[0][0], right.m[1][0], right.m[2][0]).normalize();
318    let mut right1 = Vec3::new(right.m[0][1], right.m[1][1], right.m[2][1]);
319    right1 = (right1 - scale_vec(right0, right0.dot(right1))).normalize();
320    let right2 = right0.cross(right1).normalize();
321    let mut rotation = Mat3::from_rows(
322        [
323            left_cols[0].x * right0.x + left_cols[1].x * right1.x + left_cols[2].x * right2.x,
324            left_cols[0].x * right0.y + left_cols[1].x * right1.y + left_cols[2].x * right2.y,
325            left_cols[0].x * right0.z + left_cols[1].x * right1.z + left_cols[2].x * right2.z,
326        ],
327        [
328            left_cols[0].y * right0.x + left_cols[1].y * right1.x + left_cols[2].y * right2.x,
329            left_cols[0].y * right0.y + left_cols[1].y * right1.y + left_cols[2].y * right2.y,
330            left_cols[0].y * right0.z + left_cols[1].y * right1.z + left_cols[2].y * right2.z,
331        ],
332        [
333            left_cols[0].z * right0.x + left_cols[1].z * right1.x + left_cols[2].z * right2.x,
334            left_cols[0].z * right0.y + left_cols[1].z * right1.y + left_cols[2].z * right2.y,
335            left_cols[0].z * right0.z + left_cols[1].z * right1.z + left_cols[2].z * right2.z,
336        ],
337    );
338    if determinant(rotation) < 0.0 {
339        for row in &mut rotation.m {
340            for value in row {
341                *value = -*value;
342            }
343        }
344    }
345    Ok((rotation, if determinant(matrix) < 0.0 { -scale } else { scale }))
346}
347
348fn exp_so3(omega: Vec3<f64>) -> Mat3<f64> {
349    let theta = omega.length();
350    if theta < 1e-12 {
351        return Mat3::from_rows(
352            [1.0, -omega.z, omega.y],
353            [omega.z, 1.0, -omega.x],
354            [-omega.y, omega.x, 1.0],
355        );
356    }
357    let axis = omega.normalize();
358    let skew =
359        Mat3::from_rows([0.0, -axis.z, axis.y], [axis.z, 0.0, -axis.x], [-axis.y, axis.x, 0.0]);
360    let skew2 = skew.mul_mat3(skew);
361    let mut result = Mat3::<f64>::identity();
362    let s = theta.sin();
363    let c = 1.0 - theta.cos();
364    for row in 0..3 {
365        for column in 0..3 {
366            result.m[row][column] += s * skew.m[row][column] + c * skew2.m[row][column];
367        }
368    }
369    result
370}
371
372#[allow(clippy::needless_range_loop)]
373fn smallest_symmetric_eigenvector(mut matrix: Vec<Vec<f64>>) -> Option<[f64; 12]> {
374    let size = matrix.len();
375    if size != 12 || matrix.iter().any(|row| row.len() != size) {
376        return None;
377    }
378    let mut vectors = vec![vec![0.0; size]; size];
379    for (index, row) in vectors.iter_mut().enumerate() {
380        row[index] = 1.0;
381    }
382    for _ in 0..size * size * 32 {
383        let mut pivot = (0, 1);
384        let mut maximum = 0.0_f64;
385        for (row, values) in matrix.iter().enumerate() {
386            for (column, &value) in values.iter().enumerate().skip(row + 1) {
387                if value.abs() > maximum {
388                    maximum = value.abs();
389                    pivot = (row, column);
390                }
391            }
392        }
393        if maximum < 1e-12 {
394            break;
395        }
396        let (p, q) = pivot;
397        let angle = 0.5 * (2.0 * matrix[p][q]).atan2(matrix[q][q] - matrix[p][p]);
398        let (sine, cosine) = angle.sin_cos();
399        for row in 0..size {
400            if row != p && row != q {
401                let rp = matrix[row][p];
402                let rq = matrix[row][q];
403                matrix[row][p] = cosine * rp - sine * rq;
404                matrix[p][row] = matrix[row][p];
405                matrix[row][q] = sine * rp + cosine * rq;
406                matrix[q][row] = matrix[row][q];
407            }
408        }
409        let pp = matrix[p][p];
410        let qq = matrix[q][q];
411        let pq = matrix[p][q];
412        matrix[p][p] = cosine * cosine * pp - 2.0 * sine * cosine * pq + sine * sine * qq;
413        matrix[q][q] = sine * sine * pp + 2.0 * sine * cosine * pq + cosine * cosine * qq;
414        matrix[p][q] = 0.0;
415        matrix[q][p] = 0.0;
416        for row in &mut vectors {
417            let rp = row[p];
418            let rq = row[q];
419            row[p] = cosine * rp - sine * rq;
420            row[q] = sine * rp + cosine * rq;
421        }
422    }
423    let index =
424        (0..size).min_by(|&left, &right| matrix[left][left].total_cmp(&matrix[right][right]))?;
425    let mut vector = vectors.iter().map(|row| row[index]).collect::<Vec<_>>();
426    let norm = vector.iter().map(|value| value * value).sum::<f64>().sqrt();
427    if !norm.is_finite() || norm <= f64::EPSILON {
428        return None;
429    }
430    for value in &mut vector {
431        *value /= norm;
432    }
433    let mut out = [0.0; 12];
434    out.copy_from_slice(&vector);
435    Some(out)
436}
437
438fn sample_unique(rng: &mut XorShift64, upper: usize, count: usize) -> Vec<usize> {
439    let mut selected = Vec::with_capacity(count);
440    while selected.len() < count {
441        let value = rng.next_usize(upper);
442        if !selected.contains(&value) {
443            selected.push(value);
444        }
445    }
446    selected
447}
448
449fn scale_vec(vector: Vec3<f64>, scale: f64) -> Vec3<f64> {
450    Vec3::new(vector.x * scale, vector.y * scale, vector.z * scale)
451}
452
453fn determinant(matrix: Mat3<f64>) -> f64 {
454    matrix.m[0][0] * (matrix.m[1][1] * matrix.m[2][2] - matrix.m[1][2] * matrix.m[2][1])
455        - matrix.m[0][1] * (matrix.m[1][0] * matrix.m[2][2] - matrix.m[1][2] * matrix.m[2][0])
456        + matrix.m[0][2] * (matrix.m[1][0] * matrix.m[2][1] - matrix.m[1][1] * matrix.m[2][0])
457}
458
459struct XorShift64 {
460    state: u64,
461}
462
463impl XorShift64 {
464    fn new(seed: u64) -> Self {
465        Self { state: seed | 1 }
466    }
467
468    fn next_u64(&mut self) -> u64 {
469        self.state ^= self.state << 13;
470        self.state ^= self.state >> 7;
471        self.state ^= self.state << 17;
472        self.state
473    }
474
475    fn next_usize(&mut self, upper: usize) -> usize {
476        (self.next_u64() as usize) % upper.max(1)
477    }
478}
479
480#[cfg(test)]
481mod tests {
482    use super::{project_object_point, solve_pnp, solve_pnp_ransac};
483    use crate::{AbsolutePose, CameraMatrix3, ObjectImageCorrespondence, RobustEstimationOptions};
484    use spatialrust_camera::CameraIntrinsics;
485    use spatialrust_math::{Mat3, Vec2, Vec3};
486
487    fn camera() -> CameraMatrix3 {
488        let intrinsics = CameraIntrinsics::try_new(500.0, 500.0, 320.0, 240.0, 640, 480).unwrap();
489        CameraMatrix3::from_intrinsics(intrinsics)
490    }
491
492    fn sample_pose() -> AbsolutePose {
493        AbsolutePose::try_new(
494            Mat3::from_rows(
495                [0.936_293_4, -0.275_095_9, 0.218_350_8],
496                [0.289_629_5, 0.956_425_1, -0.036_957_0],
497                [-0.198_669_3, 0.097_843_4, 0.975_170_3],
498            ),
499            Vec3::new(0.15, -0.05, 2.5),
500        )
501        .unwrap()
502    }
503
504    #[test]
505    fn solve_pnp_recovers_known_pose() {
506        let camera = camera();
507        let pose = sample_pose();
508        let objects = [
509            Vec3::new(0.0, 0.0, 0.0),
510            Vec3::new(0.4, 0.0, 0.0),
511            Vec3::new(0.0, 0.3, 0.0),
512            Vec3::new(0.0, 0.0, 0.2),
513            Vec3::new(0.25, 0.2, 0.1),
514            Vec3::new(-0.1, 0.15, -0.05),
515            Vec3::new(0.1, -0.2, 0.05),
516            Vec3::new(-0.2, -0.1, 0.15),
517        ];
518        let pairs = objects
519            .into_iter()
520            .map(|object| {
521                let image = project_object_point(pose, camera, object).unwrap();
522                ObjectImageCorrespondence::try_new(object, image).unwrap()
523            })
524            .collect::<Vec<_>>();
525        let estimated = solve_pnp(&pairs, camera).unwrap();
526        for (expected, actual) in
527            pose.rotation().m.iter().flatten().zip(estimated.rotation().m.iter().flatten())
528        {
529            assert!((expected - actual).abs() < 2e-3);
530        }
531        assert!((pose.translation().x - estimated.translation().x).abs() < 2e-3);
532        assert!((pose.translation().y - estimated.translation().y).abs() < 2e-3);
533        assert!((pose.translation().z - estimated.translation().z).abs() < 2e-3);
534    }
535
536    #[test]
537    fn solve_pnp_ransac_rejects_outliers() {
538        let camera = camera();
539        let pose = sample_pose();
540        let mut pairs = (0..20)
541            .map(|index| {
542                let object = Vec3::new(
543                    (index % 5) as f64 * 0.1 - 0.2,
544                    (index / 5) as f64 * 0.1 - 0.15,
545                    (index % 3) as f64 * 0.05,
546                );
547                let image = project_object_point(pose, camera, object).unwrap();
548                ObjectImageCorrespondence::try_new(object, image).unwrap()
549            })
550            .collect::<Vec<_>>();
551        for pair in pairs.iter_mut().take(4) {
552            *pair = ObjectImageCorrespondence::try_new(
553                pair.object(),
554                Vec2 { x: pair.image().x + 80.0, y: pair.image().y - 60.0 },
555            )
556            .unwrap();
557        }
558        let estimate = solve_pnp_ransac(
559            &pairs,
560            camera,
561            RobustEstimationOptions {
562                threshold: 2.0,
563                confidence: 0.99,
564                max_iterations: 500,
565                seed: 7,
566            },
567        )
568        .unwrap();
569        assert!(estimate.inlier_count() >= 14);
570        let recovered = *estimate.model();
571        assert!((pose.translation().z - recovered.translation().z).abs() < 5e-2);
572    }
573}