use numeris::Matrix3;
use tracing::debug;
use crate::camera_model::CameraModel;
use crate::centroid::Centroid;
use crate::distortion::fit::{
fit_polynomial_distortion, fit_radial_distortion, DistortionFitConfig,
};
use crate::solver::wcs_refine;
use crate::solver::{focal_length_from_fov, pixel_scale_from_fov, SolveResult, SolverDatabase};
use super::fit::{
build_id_lookup, compute_corrected_rmse, fit_polynomial_sigma_clip,
fit_radial_centered_sigma_clip, intrinsics_residuals, masked_rms, matched_pairs,
project_to_matched_point, MatchedPoint, MIN_RADIAL_POINTS,
};
use super::polynomial::{num_coeffs, PolynomialDistortion};
use super::Distortion;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DistortionModelType {
Polynomial { order: u32 },
Radial,
}
impl Default for DistortionModelType {
fn default() -> Self {
DistortionModelType::Polynomial { order: 4 }
}
}
#[derive(Debug, Clone)]
pub struct CalibrateConfig {
pub model: DistortionModelType,
pub max_iterations: u32,
pub sigma_clip: f64,
pub convergence_threshold_px: f64,
}
impl Default for CalibrateConfig {
fn default() -> Self {
Self {
model: DistortionModelType::default(),
max_iterations: 20,
sigma_clip: 3.0,
convergence_threshold_px: 0.01,
}
}
}
impl From<&CalibrateConfig> for DistortionFitConfig {
fn from(config: &CalibrateConfig) -> Self {
Self {
sigma_clip: config.sigma_clip,
max_iterations: config.max_iterations,
..Default::default()
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct CalibrateResult {
pub camera_model: CameraModel,
pub rmse_before_px: f64,
pub rmse_after_px: f64,
pub n_inliers: usize,
pub n_outliers: usize,
pub iterations: u32,
}
pub fn calibrate_camera(
solve_results: &[&SolveResult],
centroids: &[&[Centroid]],
database: &SolverDatabase,
image_width: u32,
image_height: u32,
config: &CalibrateConfig,
) -> crate::Result<CalibrateResult> {
assert_eq!(
solve_results.len(),
centroids.len(),
"solve_results and centroids must have the same length"
);
if let DistortionModelType::Polynomial { order } = config.model {
assert!(
(2..=6).contains(&order),
"polynomial order must be in [2, 6]"
);
}
let n_valid = solve_results.iter().filter(|sr| sr.is_ok()).count();
if n_valid == 0 {
return Err(crate::Error::InvalidInput(
"calibrate_camera: no successful solves in input".into(),
));
}
if n_valid == 1 {
single_image_calibrate(
solve_results,
centroids,
database,
image_width,
image_height,
config,
)
} else {
multi_image_calibrate(
solve_results,
centroids,
database,
image_width,
image_height,
config,
)
}
}
fn extract_crpix(distortion: Distortion) -> ([f64; 2], Distortion) {
match distortion {
Distortion::Polynomial(poly) => {
let crpix_x = poly.a_coeffs[0] * poly.scale;
let crpix_y = poly.b_coeffs[0] * poly.scale;
let mut a = poly.a_coeffs.clone();
let mut b = poly.b_coeffs.clone();
a[0] = 0.0;
b[0] = 0.0;
let new_poly = PolynomialDistortion::new(poly.order, poly.scale, a, b);
([crpix_x, crpix_y], Distortion::Polynomial(new_poly))
}
other => ([0.0, 0.0], other),
}
}
fn single_image_calibrate(
solve_results: &[&SolveResult],
centroids: &[&[Centroid]],
database: &SolverDatabase,
image_width: u32,
image_height: u32,
config: &CalibrateConfig,
) -> crate::Result<CalibrateResult> {
let fit_config = DistortionFitConfig::from(config);
let fit_result = match config.model {
DistortionModelType::Polynomial { order } => fit_polynomial_distortion(
solve_results,
centroids,
database,
image_width,
order,
&fit_config,
),
DistortionModelType::Radial => {
fit_radial_distortion(solve_results, centroids, database, image_width, &fit_config)
}
};
let first_solution = solve_results
.iter()
.find_map(|sr| sr.as_ref().ok())
.ok_or_else(|| {
crate::Error::InvalidInput("calibrate_camera: no successful solves in input".into())
})?;
let fov_rad = first_solution.fov_rad;
let parity_flip = first_solution.parity_flip;
let f_anchor = focal_length_from_fov(image_width, fov_rad as f64);
let focal_length_px = f_anchor * fit_result.focal_scale;
let (crpix, distortion) = extract_crpix(fit_result.model);
let cam = CameraModel {
focal_length_px,
image_width,
image_height,
crpix,
parity_flip,
distortion,
};
debug!(
"calibrate_camera (single, {:?}): crpix=[{:.2}, {:.2}], RMSE {:.3} -> {:.3} px, {}/{} inliers",
config.model,
crpix[0], crpix[1],
fit_result.rmse_before_px,
fit_result.rmse_after_px,
fit_result.n_inliers,
fit_result.n_inliers + fit_result.n_outliers,
);
Ok(CalibrateResult {
camera_model: cam,
rmse_before_px: fit_result.rmse_before_px,
rmse_after_px: fit_result.rmse_after_px,
n_inliers: fit_result.n_inliers,
n_outliers: fit_result.n_outliers,
iterations: fit_result.iterations,
})
}
fn multi_image_calibrate(
solve_results: &[&SolveResult],
centroids: &[&[Centroid]],
database: &SolverDatabase,
image_width: u32,
image_height: u32,
config: &CalibrateConfig,
) -> crate::Result<CalibrateResult> {
let scale = image_width as f64 / 2.0;
let id_to_idx = build_id_lookup(database);
let n_valid = solve_results.iter().filter(|sr| sr.is_ok()).count();
let n_flipped = solve_results
.iter()
.copied()
.flatten()
.filter(|sol| sol.parity_flip)
.count();
let parity_flip = 2 * n_flipped > n_valid;
let mut fovs: Vec<f32> = solve_results
.iter()
.copied()
.flatten()
.filter(|sol| sol.parity_flip == parity_flip)
.map(|sol| sol.fov_rad)
.collect();
fovs.sort_by(f32::total_cmp);
if fovs.is_empty() {
return Err(crate::Error::InvalidInput(
"calibrate_camera: no solves agree on parity".into(),
));
}
let median_fov = fovs[fovs.len() / 2];
let mut global_pixel_scale = pixel_scale_from_fov(image_width, median_fov as f64);
let parity_sign: f64 = if parity_flip { -1.0 } else { 1.0 };
debug!(
"calibrate_camera (multi): {} valid images, median FOV={:.3} deg, parity={}",
fovs.len(),
median_fov.to_degrees(),
parity_flip,
);
let mut current_distortion = Distortion::None;
let mut scale_correction = 1.0_f64;
let mut last_rmse = f64::MAX;
let mut last_rmse_before = 0.0_f64;
let fit_config = DistortionFitConfig::from(config);
struct ImageData<'a> {
idx: usize,
sol: &'a crate::solver::Solution,
centroids: &'a [Centroid],
rotation: Matrix3<f32>,
fov_rad: f32,
}
let mut image_data: Vec<ImageData> = Vec::new();
for (idx, sr) in solve_results.iter().enumerate() {
let Ok(sol) = sr else {
continue;
};
if sol.parity_flip != parity_flip {
debug!(
"calibrate_camera (multi): image {} parity_flip={} disagrees with consensus {} — \
likely a false mirror-image solve; excluding from calibration",
idx, sol.parity_flip, parity_flip,
);
continue;
}
image_data.push(ImageData {
idx,
sol,
centroids: centroids[idx],
rotation: sol.qicrs2cam.to_rotation_matrix(),
fov_rad: sol.fov_rad,
});
}
let mut total_iterations = 0u32;
let mut final_mask = Vec::new();
let mut final_n_points = 0usize;
for outer in 0..3 {
struct RefinedImage<'a> {
centroids: &'a [Centroid],
matches: Vec<(usize, usize)>, crval_ra: f64,
crval_dec: f64,
cd_matrix: [[f64; 2]; 2],
}
let mut refined_images: Vec<RefinedImage> = Vec::new();
for img in &image_data {
let cents = img.centroids;
let per_image_ps = {
let f = focal_length_from_fov(image_width, img.fov_rad as f64);
1.0 / (f * scale_correction)
};
let centroids_px: Vec<(f64, f64)> = cents
.iter()
.map(|c| {
let (ux, uy) = current_distortion.undistort(c.x as f64, c.y as f64);
(parity_sign * ux, uy)
})
.collect();
let initial_matches: Vec<(usize, usize)> =
matched_pairs(img.sol, cents.len(), &id_to_idx).collect();
if initial_matches.len() < 4 {
continue;
}
let match_radius_rad = 0.01 * img.fov_rad;
let wcs_result = wcs_refine::wcs_refine(
&img.rotation,
&initial_matches,
¢roids_px,
&database.star_vectors,
&database.star_catalog,
per_image_ps,
parity_flip,
match_radius_rad,
cents.len().min(500),
10,
);
if wcs_result.matches.len() < 4 {
debug!(
" multi-cal outer {}: image {} wcs_refine returned only {} matches, skipping",
outer,
img.idx,
wcs_result.matches.len()
);
continue;
}
debug!(
" multi-cal outer {}: image {} refined: {} matches, RMSE={:.2}\"",
outer,
img.idx,
wcs_result.matches.len(),
wcs_result.rmse_rad.to_degrees() * 3600.0,
);
refined_images.push(RefinedImage {
centroids: cents,
matches: wcs_result.matches,
crval_ra: wcs_result.crval_rad[0],
crval_dec: wcs_result.crval_rad[1],
cd_matrix: wcs_result.cd_matrix,
});
}
if refined_images.is_empty() {
debug!(" multi-cal outer {}: no refined images, aborting", outer);
break;
}
let mut all_points: Vec<MatchedPoint> = Vec::new();
for ref_img in &refined_images {
let cents = ref_img.centroids;
let (rot, _fov, _parity) = wcs_refine::wcs_to_rotation(
&ref_img.cd_matrix,
ref_img.crval_ra,
ref_img.crval_dec,
image_width,
);
for &(cent_idx, cat_idx) in &ref_img.matches {
let sv = &database.star_vectors[cat_idx];
let x_obs = cents[cent_idx].x as f64;
let y_obs = cents[cent_idx].y as f64;
if let Some(mp) =
project_to_matched_point(rot, sv, parity_sign, global_pixel_scale, x_obs, y_obs)
{
all_points.push(mp);
}
}
}
let min_points = match config.model {
DistortionModelType::Polynomial { order } => num_coeffs(order),
DistortionModelType::Radial => MIN_RADIAL_POINTS,
};
if all_points.len() < min_points {
debug!(
" multi-cal outer {}: too few points ({}) for {:?} fit",
outer,
all_points.len(),
config.model,
);
break;
}
debug!(
" multi-cal outer {}: {} total matched points from {} images",
outer,
all_points.len(),
refined_images.len(),
);
let (dist, mask, iters, rmse_after) = match config.model {
DistortionModelType::Polynomial { order } => {
let fit = fit_polynomial_sigma_clip(&all_points, order, scale, &fit_config);
let model = PolynomialDistortion::new(order, scale, fit.a_coeffs, fit.b_coeffs);
let dist = Distortion::Polynomial(model);
let rmse_after = compute_corrected_rmse(&all_points, &fit.mask, &dist);
(dist, fit.mask, fit.iterations, rmse_after)
}
DistortionModelType::Radial => {
let fit = fit_radial_centered_sigma_clip(&all_points, &fit_config);
let residuals = intrinsics_residuals(
&all_points,
&[
fit.cx, fit.cy, fit.gamma, fit.k1, fit.k2, fit.k3, fit.p1, fit.p2,
],
);
let rmse_after = masked_rms(&residuals, &fit.mask);
global_pixel_scale /= fit.gamma;
scale_correction *= fit.gamma;
debug!(
" multi-cal outer {}: radial fit gamma={:.6}, cx={:.1}, cy={:.1} folded into focal length",
outer, fit.gamma, fit.cx, fit.cy,
);
(
Distortion::Radial(fit.rescaled_model()),
fit.mask,
fit.iterations,
rmse_after,
)
}
};
let n_inliers = mask.iter().filter(|&&m| m).count();
let rmse_before = compute_corrected_rmse(&all_points, &mask, &Distortion::None);
debug!(
" multi-cal outer {}: {:?} fit: {}/{} inliers, RMSE {:.3} -> {:.3} px",
outer,
config.model,
n_inliers,
all_points.len(),
rmse_before,
rmse_after,
);
total_iterations += iters;
final_mask = mask;
final_n_points = all_points.len();
current_distortion = dist;
last_rmse_before = rmse_before;
const RMSE_REL_CONVERGENCE: f64 = 0.01; let rmse_change = (last_rmse - rmse_after).abs();
let rmse_frac_change = if last_rmse > 1e-12 {
rmse_change / last_rmse
} else {
0.0
};
last_rmse = rmse_after;
if rmse_frac_change < RMSE_REL_CONVERGENCE || rmse_change < config.convergence_threshold_px
{
debug!(
" multi-cal: converged at outer iteration {} (RMSE change={:.4} px, {:.2}%)",
outer,
rmse_change,
rmse_frac_change * 100.0,
);
break;
}
}
let (crpix, distortion) = match current_distortion {
Distortion::Polynomial(_) => extract_crpix(current_distortion),
Distortion::None | Distortion::Radial(_) => ([0.0, 0.0], current_distortion),
};
let cam = CameraModel {
focal_length_px: 1.0 / global_pixel_scale,
image_width,
image_height,
crpix,
parity_flip,
distortion,
};
if !last_rmse.is_finite() {
return Err(crate::Error::InvalidInput(
"calibrate_camera: no distortion fit completed (too few matched points)".into(),
));
}
let n_inliers = final_mask.iter().filter(|&&m| m).count();
debug!(
"calibrate_camera (multi, {:?}): crpix=[{:.2}, {:.2}], RMSE {:.3} -> {:.3} px, {}/{} inliers",
config.model, crpix[0], crpix[1], last_rmse_before, last_rmse, n_inliers, final_n_points,
);
Ok(CalibrateResult {
camera_model: cam,
rmse_before_px: last_rmse_before,
rmse_after_px: last_rmse,
n_inliers,
n_outliers: final_n_points - n_inliers,
iterations: total_iterations,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_calibrate_config_defaults() {
let cfg = CalibrateConfig::default();
assert!(matches!(
cfg.model,
DistortionModelType::Polynomial { order: 4 }
));
assert_eq!(cfg.max_iterations, 20);
assert!((cfg.sigma_clip - 3.0).abs() < 1e-12);
assert!((cfg.convergence_threshold_px - 0.01).abs() < 1e-12);
}
}