use crate::ransac::{Consensus, Estimator, RansacConfig, RansacResult, Sampler};
use rayon::prelude::*;
type ChunkBest<M> = Option<(M, f64, usize, Vec<bool>)>;
pub fn run<E, C, S>(
estimator: &E,
consensus: &C,
sampler: &mut S,
samples: &[E::Sample],
cfg: &RansacConfig,
) -> RansacResult<E::Model>
where
E: Estimator,
E::Sample: Copy,
E::Model: Clone,
C: Consensus,
S: Sampler,
{
let n = samples.len();
if n < E::SAMPLE_SIZE || cfg.max_iters == 0 {
return RansacResult {
model: None,
inliers: Vec::new(),
num_iters: 0,
score: f64::NEG_INFINITY,
};
}
let mut residuals = vec![0.0f64; n];
let mut current_inliers: Vec<bool> = Vec::with_capacity(n);
let mut best_inliers: Vec<bool> = Vec::with_capacity(n);
let mut sample_idx = vec![0usize; E::SAMPLE_SIZE];
let mut sample_buf: Vec<E::Sample> = Vec::with_capacity(E::SAMPLE_SIZE);
let mut lo_inlier_buf: Vec<E::Sample> = Vec::new();
let mut lo_models: Vec<E::Model> = Vec::new();
let mut accepted_since_lo: u32 = 0;
let mut models: Vec<E::Model> = Vec::with_capacity(10);
let mut best_score = f64::NEG_INFINITY;
let mut best_model: Option<E::Model> = None;
let mut max_iters = cfg.max_iters;
let mut i: u32 = 0;
while i < max_iters {
sampler.sample(n, &mut sample_idx);
sample_buf.clear();
for &idx in sample_idx.iter() {
sample_buf.push(samples[idx]);
}
models.clear();
estimator.fit(&sample_buf, &mut models);
for model in models.iter() {
estimator.residual_batch(model, samples, &mut residuals);
let outcome = consensus.consensus(&residuals, &mut current_inliers);
if outcome.score > best_score {
best_score = outcome.score;
best_model = Some(model.clone());
std::mem::swap(&mut best_inliers, &mut current_inliers);
accepted_since_lo += 1;
if outcome.inlier_count > 0 {
let w = outcome.inlier_count as f64 / n as f64;
let new_max = adaptive_max_iters(w, E::SAMPLE_SIZE, cfg.confidence, max_iters);
if new_max < max_iters {
max_iters = new_max;
}
}
}
}
if cfg.lo_every > 0 && accepted_since_lo >= cfg.lo_every && best_inliers.iter().any(|&b| b)
{
accepted_since_lo = 0;
lo_inlier_buf.clear();
for (idx, &is_in) in best_inliers.iter().enumerate() {
if is_in {
lo_inlier_buf.push(samples[idx]);
}
}
if lo_inlier_buf.len() > E::SAMPLE_SIZE {
lo_models.clear();
estimator.refit(&lo_inlier_buf, &mut lo_models);
for lo_model in lo_models.iter() {
estimator.residual_batch(lo_model, samples, &mut residuals);
let lo_outcome = consensus.consensus(&residuals, &mut current_inliers);
if lo_outcome.score > best_score {
best_score = lo_outcome.score;
best_model = Some(lo_model.clone());
std::mem::swap(&mut best_inliers, &mut current_inliers);
if lo_outcome.inlier_count > 0 {
let w = lo_outcome.inlier_count as f64 / n as f64;
let new_max =
adaptive_max_iters(w, E::SAMPLE_SIZE, cfg.confidence, max_iters);
if new_max < max_iters {
max_iters = new_max;
}
}
}
}
}
}
i += 1;
}
RansacResult {
model: best_model,
inliers: best_inliers,
num_iters: i,
score: best_score,
}
}
pub fn run_parallel<E, C, S>(
estimator: &E,
consensus: &C,
sampler: &mut S,
samples: &[E::Sample],
cfg: &RansacConfig,
) -> RansacResult<E::Model>
where
E: Estimator + Sync,
E::Sample: Copy + Send + Sync,
E::Model: Clone + Send + Sync,
C: Consensus + Sync,
S: Sampler,
{
let n = samples.len();
if n < E::SAMPLE_SIZE || cfg.max_iters == 0 {
return RansacResult {
model: None,
inliers: Vec::new(),
num_iters: 0,
score: f64::NEG_INFINITY,
};
}
let n_threads = rayon::current_num_threads().max(1);
let chunk_size = (n_threads * 4).max(32);
let mut chunk_samples: Vec<Vec<E::Sample>> = Vec::with_capacity(chunk_size);
let mut sample_idx = vec![0usize; E::SAMPLE_SIZE];
let mut best_score = f64::NEG_INFINITY;
let mut best_model: Option<E::Model> = None;
let mut best_inliers: Vec<bool> = Vec::with_capacity(n);
let mut max_iters = cfg.max_iters;
let mut i: u32 = 0;
while i < max_iters {
let remaining = max_iters - i;
let this_chunk = chunk_size.min(remaining as usize);
chunk_samples.clear();
for _ in 0..this_chunk {
sampler.sample(n, &mut sample_idx);
let mut buf: Vec<E::Sample> = Vec::with_capacity(E::SAMPLE_SIZE);
for &idx in sample_idx.iter() {
buf.push(samples[idx]);
}
chunk_samples.push(buf);
}
let chunk_results: Vec<ChunkBest<E::Model>> = chunk_samples
.par_iter()
.map(|sample_buf| {
let mut models: Vec<E::Model> = Vec::with_capacity(10);
let mut residuals = vec![0.0f64; n];
let mut inliers: Vec<bool> = Vec::with_capacity(n);
estimator.fit(sample_buf, &mut models);
let mut local_best: Option<(E::Model, f64, usize, Vec<bool>)> = None;
for model in models.iter() {
estimator.residual_batch(model, samples, &mut residuals);
let outcome = consensus.consensus(&residuals, &mut inliers);
let take = match &local_best {
None => true,
Some((_, s, _, _)) => outcome.score > *s,
};
if take {
local_best = Some((
model.clone(),
outcome.score,
outcome.inlier_count,
inliers.clone(),
));
}
}
local_best
})
.collect();
for entry in chunk_results.into_iter().flatten() {
let (model, score, inlier_count, inliers) = entry;
if score > best_score {
best_score = score;
best_model = Some(model);
best_inliers = inliers;
if inlier_count > 0 {
let w = inlier_count as f64 / n as f64;
let new_max = adaptive_max_iters(w, E::SAMPLE_SIZE, cfg.confidence, max_iters);
if new_max < max_iters {
max_iters = new_max;
}
}
}
}
i += this_chunk as u32;
}
RansacResult {
model: best_model,
inliers: best_inliers,
num_iters: i,
score: best_score,
}
}
#[inline]
fn adaptive_max_iters(inlier_ratio: f64, sample_size: usize, confidence: f64, current: u32) -> u32 {
if inlier_ratio <= 0.0 {
return current;
}
let p_all_inlier = inlier_ratio.powi(sample_size as i32);
if p_all_inlier >= 1.0 {
return 1;
}
let conf = confidence.clamp(0.0, 1.0 - 1e-12);
let denom = (1.0 - p_all_inlier).ln();
if denom >= 0.0 || !denom.is_finite() {
return current;
}
let raw = (1.0 - conf).ln() / denom;
if !raw.is_finite() || raw <= 0.0 {
return current;
}
let ceiled = raw.ceil() as u32;
ceiled.min(current).max(1)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ransac::{
estimators::FundamentalEstimator, Match2d2d, ThresholdConsensus, UniformSampler,
};
use kornia_algebra::{Vec2F64, Vec3F64};
use rand::{rngs::StdRng, SeedableRng};
#[test]
fn under_min_samples_returns_empty() {
let est = FundamentalEstimator;
let consensus = ThresholdConsensus { threshold: 1.0 };
let mut sampler = UniformSampler::new(StdRng::seed_from_u64(0));
let result: RansacResult<_> = run(
&est,
&consensus,
&mut sampler,
&[Match2d2d::new(Vec2F64::new(0.0, 0.0), Vec2F64::new(0.0, 0.0)); 5],
&RansacConfig::default(),
);
assert!(result.model.is_none());
assert_eq!(result.num_iters, 0);
assert!(result.inliers.is_empty());
}
#[test]
fn recovers_inliers_under_outliers() {
let pair = synthetic_with_outliers(60, 40, 12345);
let est = FundamentalEstimator;
let consensus = ThresholdConsensus { threshold: 4.0 }; let mut sampler = UniformSampler::new(StdRng::seed_from_u64(0xC0FFEE));
let cfg = RansacConfig {
max_iters: 1000,
confidence: 0.999,
inlier_threshold: 4.0,
..Default::default()
};
let result = run(&est, &consensus, &mut sampler, &pair.matches, &cfg);
assert!(result.model.is_some(), "driver returned no model");
let recovered = result.inlier_count();
assert!(
recovered >= 48,
"recovered only {recovered}/60 true inliers (score = {})",
result.score
);
assert!(
result.num_iters < cfg.max_iters,
"adaptive cap never engaged: ran {} of {} iters",
result.num_iters,
cfg.max_iters,
);
}
#[test]
fn run_parallel_matches_serial_quality() {
let pair = synthetic_with_outliers(80, 50, 0xCAFE);
let est = FundamentalEstimator;
let consensus = ThresholdConsensus { threshold: 4.0 };
let cfg = RansacConfig {
max_iters: 600,
confidence: 0.999,
inlier_threshold: 4.0,
..Default::default()
};
let mut sampler_serial = UniformSampler::new(StdRng::seed_from_u64(11));
let serial = run(&est, &consensus, &mut sampler_serial, &pair.matches, &cfg);
let mut sampler_par = UniformSampler::new(StdRng::seed_from_u64(11));
let par = run_parallel(&est, &consensus, &mut sampler_par, &pair.matches, &cfg);
assert!(serial.model.is_some() && par.model.is_some());
let serial_inliers = serial.inliers[..80].iter().filter(|&&b| b).count();
let par_inliers = par.inliers[..80].iter().filter(|&&b| b).count();
assert!(serial_inliers >= 64, "serial: {serial_inliers}");
assert!(par_inliers >= 64, "parallel: {par_inliers}");
let ratio = par.score / serial.score;
assert!(
ratio > 0.85,
"parallel score {} regressed vs serial {}",
par.score,
serial.score
);
}
#[test]
fn lo_ransac_does_not_regress_plain_ransac() {
let pair = synthetic_with_outliers(60, 40, 0xBEEF);
let est = FundamentalEstimator;
let consensus = ThresholdConsensus { threshold: 4.0 };
let cfg_plain = RansacConfig {
max_iters: 500,
confidence: 0.999,
inlier_threshold: 4.0,
lo_every: 0,
..Default::default()
};
let cfg_lo = RansacConfig {
lo_every: 5,
..cfg_plain.clone()
};
let mut sampler_plain = UniformSampler::new(StdRng::seed_from_u64(7));
let plain = run(
&est,
&consensus,
&mut sampler_plain,
&pair.matches,
&cfg_plain,
);
let mut sampler_lo = UniformSampler::new(StdRng::seed_from_u64(7));
let lo = run(&est, &consensus, &mut sampler_lo, &pair.matches, &cfg_lo);
assert!(
lo.score >= plain.score,
"LO regressed: plain.score={} lo.score={}",
plain.score,
lo.score
);
assert!(lo.model.is_some());
}
struct Pair {
matches: Vec<Match2d2d>,
}
fn synthetic_with_outliers(n_inliers: usize, n_outliers: usize, seed: u64) -> Pair {
let fx = 500.0_f64;
let fy = 500.0_f64;
let cx = 320.0_f64;
let cy = 240.0_f64;
let angle = 0.1_f64;
let r = [
[angle.cos(), 0.0, -angle.sin()],
[0.0, 1.0, 0.0],
[angle.sin(), 0.0, angle.cos()],
];
let t = [1.0_f64, 0.0, 0.2];
let mut state = seed;
let mut next = || {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 33) as f64) / (1u64 << 31) as f64 - 1.0
};
let mut matches = Vec::with_capacity(n_inliers + n_outliers);
for _ in 0..n_inliers {
let p = Vec3F64::new(next() * 0.6, next() * 0.6, 3.0 + next().abs() * 3.0);
let u1 = fx * p.x / p.z + cx;
let v1 = fy * p.y / p.z + cy;
let pc2 = [
r[0][0] * p.x + r[0][1] * p.y + r[0][2] * p.z + t[0],
r[1][0] * p.x + r[1][1] * p.y + r[1][2] * p.z + t[1],
r[2][0] * p.x + r[2][1] * p.y + r[2][2] * p.z + t[2],
];
let u2 = fx * pc2[0] / pc2[2] + cx;
let v2 = fy * pc2[1] / pc2[2] + cy;
matches.push(Match2d2d::new(Vec2F64::new(u1, v1), Vec2F64::new(u2, v2)));
}
for _ in 0..n_outliers {
let u1 = (next() * 0.5 + 0.5) * 640.0;
let v1 = (next() * 0.5 + 0.5) * 480.0;
let u2 = (next() * 0.5 + 0.5) * 640.0;
let v2 = (next() * 0.5 + 0.5) * 480.0;
matches.push(Match2d2d::new(Vec2F64::new(u1, v1), Vec2F64::new(u2, v2)));
}
Pair { matches }
}
}