use std::cell::RefCell;
use crate::frame::Frame;
use crate::matcher::{Match, Matcher};
use crate::template::Template;
use corrmatch::{
CompileConfigNoRot, CompiledTemplate, ImageView, MatchConfig, Matcher as CorrInner,
RotationMode, Template as CorrTemplate,
};
struct Cached {
key: u64,
inner: CorrInner,
}
fn no_rot_levels(w: usize, h: usize) -> usize {
let mut levels = 0usize;
let (mut cw, mut ch) = (w, h);
while levels < 6 && cw.min(ch) >= 4 {
levels += 1;
cw /= 2;
ch /= 2;
}
levels.max(1)
}
#[derive(Clone, Debug)]
pub struct CorrConfig {
pub max_image_levels: usize,
pub beam_width: usize,
pub roi_radius: usize,
pub min_score: f32,
pub parallel: bool,
}
impl Default for CorrConfig {
fn default() -> Self {
CorrConfig {
max_image_levels: 6,
beam_width: 8,
roi_radius: 8,
min_score: f32::NEG_INFINITY,
parallel: cfg!(feature = "parallel"),
}
}
}
impl CorrConfig {
fn sanitize(&self) -> CorrConfig {
CorrConfig {
max_image_levels: self.max_image_levels.max(1),
beam_width: self.beam_width.max(1),
roi_radius: self.roi_radius.max(1),
min_score: match self.min_score {
v if v.is_nan() || v == f32::NEG_INFINITY => f32::NEG_INFINITY,
v if v == f32::INFINITY => f32::MAX,
v => v,
},
parallel: self.parallel && cfg!(feature = "parallel"),
}
}
}
pub struct CorrMatcher {
cfg: CorrConfig,
cache: RefCell<Option<Cached>>,
gray: RefCell<Vec<u8>>,
}
impl CorrMatcher {
pub fn new() -> Self {
Self::with_config(CorrConfig::default())
}
pub fn with_config(cfg: CorrConfig) -> Self {
CorrMatcher {
cfg: cfg.sanitize(),
cache: RefCell::new(None),
gray: RefCell::new(Vec::new()),
}
}
pub fn config(&self) -> &CorrConfig {
&self.cfg
}
fn ensure_compiled(&self, tpl: &Template) -> Option<()> {
let key = tpl.content_key();
let mut cache = self.cache.borrow_mut();
if cache.as_ref().map(|c| c.key) == Some(key) {
return Some(());
}
let corr_tpl = CorrTemplate::new(tpl.to_gray(), tpl.width, tpl.height).ok()?;
let compile_cfg = CompileConfigNoRot {
max_levels: no_rot_levels(tpl.width, tpl.height),
};
let compiled = CompiledTemplate::compile_unrotated(&corr_tpl, compile_cfg).ok()?;
let inner = CorrInner::new(compiled).with_config(MatchConfig {
rotation: RotationMode::Disabled,
parallel: self.cfg.parallel,
max_image_levels: self.cfg.max_image_levels,
beam_width: self.cfg.beam_width,
roi_radius: self.cfg.roi_radius,
..MatchConfig::default()
});
*cache = Some(Cached { key, inner });
Some(())
}
}
impl Default for CorrMatcher {
fn default() -> Self {
Self::new()
}
}
impl Matcher for CorrMatcher {
fn find(&self, frame: &Frame, tpl: &Template) -> Option<Match> {
self.ensure_compiled(tpl)?;
frame.to_gray_into(&mut self.gray.borrow_mut());
let cache = self.cache.borrow();
let gray = self.gray.borrow();
let inner = &cache.as_ref()?.inner;
let view = ImageView::from_slice(&gray, frame.width, frame.height).ok()?;
let res = inner.match_image(view).ok()?;
if res.score < self.cfg.min_score {
return None;
}
Some(Match {
x: res.x as i32,
y: res.y as i32,
score: res.score,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::frame::Rect;
fn noisy_bg(w: usize, h: usize) -> Vec<u8> {
let mut px = vec![0u8; w * h * 4];
for y in 0..h {
for x in 0..w {
let i = (y * w + x) * 4;
px[i] = ((x * 7 + y * 13) % 256) as u8;
px[i + 1] = ((x * 3 + y * 11) % 256) as u8;
px[i + 2] = ((x * 5 + y * 17) % 256) as u8;
px[i + 3] = 255;
}
}
px
}
fn paste_textured(px: &mut [u8], w: usize, tx: usize, ty: usize, s: usize) -> Vec<u8> {
let mut rgb = Vec::with_capacity(s * s * 3);
for yy in 0..s {
for xx in 0..s {
let i = ((ty + yy) * w + (tx + xx)) * 4;
let checker = if (xx / 4 + yy / 4) % 2 == 0 { 240 } else { 15 };
let r = (checker - xx as i32 * 2).clamp(0, 255) as u8;
let g = (checker - yy as i32 * 2).clamp(0, 255) as u8;
let b = ((xx * 5 + yy * 3) % 256) as u8;
px[i] = r;
px[i + 1] = g;
px[i + 2] = b;
rgb.extend_from_slice(&[r, g, b]);
}
}
rgb
}
#[test]
fn cached_compile_is_consistent() {
let (w, h) = (256usize, 256usize);
let mut px = noisy_bg(w, h);
let rgb = paste_textured(&mut px, w, 150, 120, 32);
let frame = Frame::rgba8(w, h, px);
let t = Template::from_rgb(rgb, 32, 32);
let m = CorrMatcher::new();
let a = m.find(&frame, &t).expect("first find");
let b = m.find(&frame, &t).expect("cached find");
assert_eq!((a.x, a.y), (150, 120));
assert_eq!(a.x, b.x);
assert_eq!(a.y, b.y);
}
#[test]
fn region_default_works() {
let (w, h) = (256usize, 256usize);
let mut px = noisy_bg(w, h);
let rgb = paste_textured(&mut px, w, 60, 70, 32);
let frame = Frame::rgba8(w, h, px);
let t = Template::from_rgb(rgb, 32, 32);
let m = CorrMatcher::new();
let hit = m
.find_in(&frame, &t, Rect::new(50, 60, 64, 64))
.expect("区域内应命中");
assert_eq!((hit.x, hit.y), (60, 70));
}
#[test]
fn levels_stop_at_non_degenerate_size() {
assert_eq!(no_rot_levels(64, 64), 5);
assert_eq!(no_rot_levels(32, 32), 4);
assert_eq!(no_rot_levels(8, 8), 2);
assert_eq!(no_rot_levels(4, 4), 1);
assert_eq!(no_rot_levels(2, 2), 1);
assert_eq!(no_rot_levels(1000, 60), 4);
}
#[test]
fn big_frame_textured_template_hits() {
let (w, h) = (960usize, 540usize);
let mut px = noisy_bg(w, h);
let rgb = paste_textured(&mut px, w, 600, 380, 64);
let frame = Frame::rgba8(w, h, px);
let t = Template::from_rgb(rgb, 64, 64);
let hit = CorrMatcher::new().find(&frame, &t).expect("全屏应命中");
assert!(
(hit.x - 600).abs() <= 2 && (hit.y - 380).abs() <= 2,
"命中偏移过多: {:?}",
(hit.x, hit.y)
);
assert!(hit.score > 0.9, "命中置信度过低: {}", hit.score);
}
#[test]
fn invalid_config_is_clamped() {
let c = CorrConfig {
max_image_levels: 0,
beam_width: 0,
roi_radius: 0,
min_score: f32::NAN,
parallel: true,
};
let s = c.sanitize();
assert_eq!(s.max_image_levels, 1);
assert_eq!(s.beam_width, 1);
assert_eq!(s.roi_radius, 1);
assert_eq!(s.min_score, f32::NEG_INFINITY);
assert_eq!(s.parallel, cfg!(feature = "parallel"));
}
#[test]
fn clamped_config_still_finds() {
let (w, h) = (256usize, 256usize);
let mut px = noisy_bg(w, h);
let rgb = paste_textured(&mut px, w, 150, 120, 32);
let frame = Frame::rgba8(w, h, px);
let t = Template::from_rgb(rgb, 32, 32);
let m = CorrMatcher::with_config(CorrConfig {
max_image_levels: 0,
beam_width: 0,
roi_radius: 0,
min_score: f32::NAN,
parallel: true,
});
let hit = m.find(&frame, &t).expect("非法入参被夹安全后应命中");
assert_eq!((hit.x, hit.y), (150, 120));
}
#[test]
fn min_score_filters_weak_hits() {
let (w, h) = (256usize, 256usize);
let mut px = noisy_bg(w, h);
let rgb = paste_textured(&mut px, w, 150, 120, 32);
let frame = Frame::rgba8(w, h, px);
let t = Template::from_rgb(rgb, 32, 32);
let loose = CorrMatcher::with_config(CorrConfig {
min_score: 0.5,
..CorrConfig::default()
});
assert!(loose.find(&frame, &t).is_some(), "高于 0.5 的命中应保留");
let strict = CorrMatcher::with_config(CorrConfig {
min_score: 1.0001,
..CorrConfig::default()
});
assert!(
strict.find(&frame, &t).is_none(),
"不可能的阈值应过滤掉结果"
);
}
}