use crate::core::error::Result;
use crate::core::geometry::Circle;
use crate::core::gradient::{GradientOperator, compute_gradient};
use crate::core::image_view::ImageView;
use crate::core::nms::NmsConfig;
use crate::core::polarity::Polarity;
use crate::core::scalar::Scalar;
use crate::diagnostics::detection::{
CircleDetectionDiagnostics, RejectedProposal, RejectionReason,
};
use crate::propose::extract::extract_proposals;
use crate::propose::frst::{FrstConfig, frst_response};
use crate::refine::circle::{CircleRefineConfig, refine_circle};
use crate::refine::result::RefinementStatus;
use crate::support::score::{
ScoringConfig, SupportScore, SupportScoreBreakdown, score_circle_support,
};
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct DetectCirclesConfig {
pub radii: Vec<u32>,
pub polarity: Polarity,
pub radius_hint: Scalar,
pub min_score: Scalar,
pub gradient_operator: GradientOperator,
pub advanced: DetectCirclesAdvanced,
}
#[derive(Debug, Clone, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub struct DetectCirclesAdvanced {
pub frst: FrstConfig,
pub nms: NmsConfig,
pub scoring: ScoringConfig,
pub refinement: CircleRefineConfig,
}
impl Default for DetectCirclesConfig {
fn default() -> Self {
Self {
radii: FrstConfig::default().radii,
polarity: Polarity::Both,
radius_hint: 10.0,
min_score: 0.0,
gradient_operator: GradientOperator::default(),
advanced: DetectCirclesAdvanced::default(),
}
}
}
impl DetectCirclesConfig {
pub fn for_radii(radii: impl IntoIterator<Item = u32>) -> Self {
Self {
radii: radii.into_iter().collect(),
..Self::default()
}
}
pub fn polarity(mut self, polarity: Polarity) -> Self {
self.polarity = polarity;
self.advanced.frst.polarity = polarity;
self
}
pub fn radius_hint(mut self, radius_hint: Scalar) -> Self {
self.radius_hint = radius_hint;
self
}
pub fn min_score(mut self, min_score: Scalar) -> Self {
self.min_score = min_score;
self
}
pub fn gradient_operator(mut self, gradient_operator: GradientOperator) -> Self {
self.gradient_operator = gradient_operator;
self
}
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(
feature = "serde",
serde(bound(
serialize = "T: serde::Serialize",
deserialize = "T: serde::de::DeserializeOwned"
))
)]
#[non_exhaustive]
pub struct Detection<T> {
pub hypothesis: T,
pub score: SupportScore,
pub status: RefinementStatus,
}
pub type CircleDetection = Detection<Circle>;
pub fn detect_circles(
image: &ImageView<'_, u8>,
config: &DetectCirclesConfig,
) -> Result<Vec<CircleDetection>> {
run_detection(image, config).map(|(detections, _diagnostics)| detections)
}
pub fn detect_circles_with_diagnostics(
image: &ImageView<'_, u8>,
config: &DetectCirclesConfig,
) -> Result<(Vec<CircleDetection>, CircleDetectionDiagnostics)> {
run_detection(image, config)
}
fn run_detection(
image: &ImageView<'_, u8>,
config: &DetectCirclesConfig,
) -> Result<(Vec<CircleDetection>, CircleDetectionDiagnostics)> {
config.advanced.refinement.validate()?;
let gradient = compute_gradient(image, config.gradient_operator)?;
let mut frst_config = config.advanced.frst.clone();
frst_config.radii = config.radii.clone();
frst_config.polarity = config.polarity;
let response = frst_response(&gradient, &frst_config)?;
let proposals = extract_proposals(&response, &config.advanced.nms, config.polarity);
let mut accepted: Vec<(CircleDetection, SupportScoreBreakdown)> = Vec::new();
let mut rejected: Vec<RejectedProposal> = Vec::new();
for proposal in &proposals {
let circle = Circle::new(proposal.seed.position, config.radius_hint);
let breakdown = score_circle_support(&gradient, &circle, &config.advanced.scoring);
if breakdown.is_degenerate {
rejected.push(RejectedProposal {
proposal: proposal.clone(),
reason: RejectionReason::Degenerate,
score: breakdown,
});
continue;
}
if breakdown.total < config.min_score {
rejected.push(RejectedProposal {
proposal: proposal.clone(),
reason: RejectionReason::LowScore,
score: breakdown,
});
continue;
}
match refine_circle(&gradient, &circle, &config.advanced.refinement) {
Ok(refined) => accepted.push((
Detection {
hypothesis: refined.hypothesis,
score: SupportScore {
total: breakdown.total,
},
status: refined.status,
},
breakdown,
)),
Err(_) => rejected.push(RejectedProposal {
proposal: proposal.clone(),
reason: RejectionReason::RefinementFailed,
score: breakdown,
}),
}
}
accepted.sort_by(|a, b| {
b.0.score
.total
.partial_cmp(&a.0.score.total)
.unwrap_or(std::cmp::Ordering::Equal)
});
let (detections, score_breakdowns): (Vec<CircleDetection>, Vec<SupportScoreBreakdown>) =
accepted.into_iter().unzip();
let diagnostics = CircleDetectionDiagnostics {
response,
proposals,
rejected,
score_breakdowns,
};
Ok((detections, diagnostics))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detect_circles_finds_synthetic_disk() {
let size = 128;
let cx = 64.0f32;
let cy = 64.0f32;
let radius = 18.0f32;
let mut data = vec![0u8; size * size];
for y in 0..size {
for x in 0..size {
let dx = x as f32 - cx;
let dy = y as f32 - cy;
if (dx * dx + dy * dy).sqrt() <= radius {
data[y * size + x] = 255;
}
}
}
let image = ImageView::from_slice(&data, size, size).unwrap();
let config = DetectCirclesConfig {
radii: vec![17, 18, 19],
polarity: Polarity::Bright,
radius_hint: radius,
advanced: DetectCirclesAdvanced {
frst: FrstConfig {
gradient_threshold: 1.0,
..FrstConfig::default()
},
..DetectCirclesAdvanced::default()
},
..DetectCirclesConfig::default()
};
let detections = detect_circles(&image, &config).unwrap();
assert!(!detections.is_empty(), "should detect the synthetic disk");
let best = &detections[0];
let dx = best.hypothesis.center.x - cx;
let dy = best.hypothesis.center.y - cy;
assert!(
(dx * dx + dy * dy).sqrt() < 3.0,
"center should be near ({cx}, {cy}), got ({}, {})",
best.hypothesis.center.x,
best.hypothesis.center.y,
);
}
#[test]
fn builder_chain_sets_expected_fields() {
let config = DetectCirclesConfig::for_radii([9, 10, 11])
.polarity(Polarity::Bright)
.radius_hint(12.5)
.min_score(0.2)
.gradient_operator(GradientOperator::Scharr);
assert_eq!(config.radii, vec![9, 10, 11]);
assert_eq!(config.polarity, Polarity::Bright);
assert_eq!(config.advanced.frst.polarity, Polarity::Bright);
assert_eq!(config.radius_hint, 12.5);
assert_eq!(config.min_score, 0.2);
assert_eq!(config.gradient_operator, GradientOperator::Scharr);
let defaults = DetectCirclesConfig::default();
assert_eq!(config.advanced.nms.radius, defaults.advanced.nms.radius);
assert_eq!(
config.advanced.refinement.max_iterations,
defaults.advanced.refinement.max_iterations
);
}
#[test]
fn invalid_refinement_config_returns_error() {
let size = 64;
let data = vec![128u8; size * size];
let image = ImageView::from_slice(&data, size, size).unwrap();
let config = DetectCirclesConfig {
advanced: DetectCirclesAdvanced {
refinement: CircleRefineConfig {
max_iterations: 0,
..CircleRefineConfig::default()
},
..DetectCirclesAdvanced::default()
},
..DetectCirclesConfig::default()
};
let result = detect_circles(&image, &config);
assert!(
matches!(
result,
Err(crate::core::error::RadSymError::InvalidConfig { .. })
),
"expected InvalidConfig error, got {result:?}"
);
}
#[test]
fn detect_circles_with_diagnostics_matches_detect_circles() {
let size = 128;
let cx = 64.0f32;
let cy = 64.0f32;
let radius = 18.0f32;
let mut data = vec![0u8; size * size];
for y in 0..size {
for x in 0..size {
let dx = x as f32 - cx;
let dy = y as f32 - cy;
if (dx * dx + dy * dy).sqrt() <= radius {
data[y * size + x] = 255;
}
}
}
let image = ImageView::from_slice(&data, size, size).unwrap();
let config = DetectCirclesConfig {
radii: vec![17, 18, 19],
polarity: Polarity::Bright,
radius_hint: radius,
advanced: DetectCirclesAdvanced {
frst: FrstConfig {
gradient_threshold: 1.0,
..FrstConfig::default()
},
..DetectCirclesAdvanced::default()
},
..DetectCirclesConfig::default()
};
let plain = detect_circles(&image, &config).unwrap();
let (detailed, diagnostics) = detect_circles_with_diagnostics(&image, &config).unwrap();
assert_eq!(plain.len(), detailed.len());
for (a, b) in plain.iter().zip(&detailed) {
assert_eq!(a.hypothesis.center.x, b.hypothesis.center.x);
assert_eq!(a.hypothesis.center.y, b.hypothesis.center.y);
assert_eq!(a.hypothesis.radius, b.hypothesis.radius);
assert_eq!(a.score.total, b.score.total);
}
assert_eq!(detailed.len(), diagnostics.score_breakdowns.len());
for (det, breakdown) in detailed.iter().zip(&diagnostics.score_breakdowns) {
assert_eq!(det.score.total, breakdown.total);
}
assert!(!diagnostics.proposals.is_empty());
}
}