use std::sync::Arc;
use crate::config::DetectionConfig;
use crate::error::Result;
use crate::inference::ModelBackend;
use crate::types::{Image, QUAD_CORNERS as REGION_CORNERS};
use super::group::Grouped;
pub(crate) trait TextDetector: Send + Sync {
fn detect(&self, input: &DetectorInput) -> Result<DetectedRegions>;
}
pub(crate) struct DetectorInput<'a> {
pub image: &'a Image,
}
pub(crate) struct DetectedRegions {
pub regions: Vec<DetectedRegion>,
}
pub(crate) struct DetectedRegion {
pub corners: [[f32; 2]; REGION_CORNERS],
pub axis_aligned: bool,
}
pub(crate) struct CraftDetector {
backend: Arc<dyn ModelBackend>,
config: DetectionConfig,
}
impl CraftDetector {
pub(crate) fn new(backend: Arc<dyn ModelBackend>, detection_config: DetectionConfig) -> Self {
Self {
backend,
config: detection_config,
}
}
}
impl TextDetector for CraftDetector {
fn detect(&self, input: &DetectorInput) -> Result<DetectedRegions> {
let prepared = super::preprocess::prepare(input.image, self.config.canvas_size, self.config.mag_ratio)?;
let heat = super::craft::run_craft(self.backend.as_ref(), prepared.tensor)?;
let mut boxes = super::postprocess::get_det_boxes(
&heat.region,
&heat.link,
self.config.text_threshold,
self.config.link_threshold,
self.config.low_text,
)?;
super::postprocess::adjust_coordinates(&mut boxes, prepared.inv_ratio);
let grouped = super::group::group_boxes(&boxes, &self.config);
let regions = map_grouped_to_regions(grouped, self.config.min_size);
Ok(DetectedRegions { regions })
}
}
fn map_grouped_to_regions(grouped: Grouped, min_size: u32) -> Vec<DetectedRegion> {
let min_size = min_size as f32;
let mut regions = Vec::with_capacity(grouped.horizontal.len() + grouped.free.len());
for [x_min, x_max, y_min, y_max] in grouped.horizontal {
let corners = [[x_min, y_min], [x_max, y_min], [x_max, y_max], [x_min, y_max]];
push_if_large_enough(&mut regions, corners, true, min_size);
}
for corners in grouped.free {
push_if_large_enough(&mut regions, corners, false, min_size);
}
regions
}
fn push_if_large_enough(
regions: &mut Vec<DetectedRegion>,
corners: [[f32; 2]; REGION_CORNERS],
axis_aligned: bool,
min_size: f32,
) {
if region_max_extent(&corners) > min_size {
regions.push(DetectedRegion { corners, axis_aligned });
}
}
fn region_max_extent(corners: &[[f32; 2]; REGION_CORNERS]) -> f32 {
let mut min_x = f32::MAX;
let mut max_x = f32::MIN;
let mut min_y = f32::MAX;
let mut max_y = f32::MIN;
for corner in corners {
min_x = min_x.min(corner[0]);
max_x = max_x.max(corner[0]);
min_y = min_y.min(corner[1]);
max_y = max_y.max(corner[1]);
}
(max_x - min_x).max(max_y - min_y)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::Tensor;
use ndarray::{ArrayD, IxDyn};
struct FixedBackend {
output: ArrayD<f32>,
}
impl ModelBackend for FixedBackend {
fn name(&self) -> &str {
"fixed"
}
fn run(&self, _input: Tensor) -> Result<Tensor> {
Ok(self.output.clone())
}
}
fn solid_image(width: u32, height: u32) -> Image {
let pixels = vec![255u8; (width * height * 3) as usize];
Image::from_rgb8(width, height, pixels).expect("valid rgb buffer")
}
#[test]
fn should_map_horizontal_box_to_axis_aligned_corners() {
let grouped = Grouped {
horizontal: vec![[10.0, 50.0, 20.0, 60.0]],
free: Vec::new(),
};
let regions = map_grouped_to_regions(grouped, 0);
assert_eq!(regions.len(), 1);
assert!(regions[0].axis_aligned);
assert_eq!(
regions[0].corners,
[[10.0, 20.0], [50.0, 20.0], [50.0, 60.0], [10.0, 60.0]]
);
}
#[test]
fn should_map_free_quad_to_non_axis_aligned_region() {
let quad = [[10.0, 12.0], [50.0, 20.0], [48.0, 60.0], [8.0, 52.0]];
let grouped = Grouped {
horizontal: Vec::new(),
free: vec![quad],
};
let regions = map_grouped_to_regions(grouped, 0);
assert_eq!(regions.len(), 1);
assert!(!regions[0].axis_aligned);
assert_eq!(regions[0].corners, quad);
}
#[test]
fn should_drop_region_smaller_than_min_size() {
let grouped = Grouped {
horizontal: vec![[10.0, 18.0, 20.0, 26.0]],
free: Vec::new(),
};
let regions = map_grouped_to_regions(grouped, 20);
assert!(regions.is_empty());
}
#[test]
fn should_drop_region_whose_extent_equals_min_size() {
let equal = Grouped {
horizontal: vec![[10.0, 30.0, 20.0, 26.0]],
free: Vec::new(),
};
assert!(map_grouped_to_regions(equal, 20).is_empty());
let above = Grouped {
horizontal: vec![[10.0, 31.0, 20.0, 26.0]],
free: Vec::new(),
};
assert_eq!(map_grouped_to_regions(above, 20).len(), 1);
}
#[test]
fn should_run_full_pipeline_and_produce_regions() {
let (height, width) = (8usize, 8usize);
let mut output = ArrayD::<f32>::zeros(IxDyn(&[1, height, width, 2]));
for row in 1..7 {
for col in 1..7 {
output[[0, row, col, 0]] = 1.0;
}
}
let backend = Arc::new(FixedBackend { output });
let config = DetectionConfig {
min_size: 0,
..DetectionConfig::default()
};
let detector = CraftDetector::new(backend, config);
let image = solid_image(16, 16);
let input = DetectorInput { image: &image };
let regions = detector.detect(&input).expect("detection succeeds");
assert!(!regions.regions.is_empty(), "expected at least one region");
}
#[cfg(feature = "ort")]
#[test]
#[ignore = "requires the ONNX Runtime native library and a CRAFT model file"]
fn detect_over_real_craft_model() {
let model_path =
std::env::var("EASYOCR_TEST_CRAFT_ONNX").expect("set EASYOCR_TEST_CRAFT_ONNX to a CRAFT ONNX model path");
let model_bytes = std::fs::read(&model_path).expect("read the model file");
let backend = crate::inference::load_backend(crate::config::Backend::Ort, &model_bytes, 1)
.expect("load the CRAFT ONNX model");
let detector = CraftDetector::new(Arc::from(backend), DetectionConfig::default());
let image = solid_image(64, 64);
let input = DetectorInput { image: &image };
let result = detector.detect(&input);
assert!(result.is_ok(), "detection over the real model must succeed");
}
}