1use 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
12pub 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
30pub 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
107pub 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 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 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 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}