use std::sync::Arc;
use ndarray::Axis;
use rayon::prelude::*;
use crate::config::{Decoder, RecognitionConfig};
use crate::error::{OcrError, Result};
use crate::inference::ModelBackend;
use crate::types::QUAD_CORNERS as REGION_CORNERS;
use super::charset::Charset;
pub(crate) trait TextRecognizer: Send + Sync {
fn recognize(&self, crops: &[RegionCrop]) -> Result<Vec<RecognizedText>>;
}
pub(crate) struct RegionCrop {
pub width: u32,
pub height: u32,
pub gray: Vec<u8>,
pub corners: [[f32; 2]; REGION_CORNERS],
}
pub(crate) struct RecognizedText {
pub text: String,
pub confidence: f32,
}
pub(crate) struct CrnnRecognizer {
backend: Arc<dyn ModelBackend>,
charset: Charset,
config: RecognitionConfig,
}
impl CrnnRecognizer {
pub(crate) fn new(backend: Arc<dyn ModelBackend>, charset: Charset, config: RecognitionConfig) -> Self {
Self {
backend,
charset,
config,
}
}
fn run_pass(&self, crops: &[RegionCrop], ignore: &[usize]) -> Result<Vec<RecognizedText>> {
let batch_size = self.config.batch_size.max(1);
let mut results = Vec::with_capacity(crops.len());
for chunk in crops.chunks(batch_size) {
let tensor = super::preprocess::prepare_batch(chunk)?;
let logits = super::crnn::run_crnn(self.backend.as_ref(), tensor)?;
let decoded: Vec<RecognizedText> = (0..chunk.len())
.into_par_iter()
.map(|row| super::ctc::decode_greedy(logits.index_axis(Axis(0), row), &self.charset, ignore))
.collect();
results.extend(decoded);
}
Ok(results)
}
fn apply_second_pass(&self, crops: &[RegionCrop], ignore: &[usize], results: &mut [RecognizedText]) -> Result<()> {
let indices: Vec<usize> = results
.iter()
.enumerate()
.filter(|(_, result)| result.confidence < self.config.contrast_ths)
.map(|(index, _)| index)
.collect();
if indices.is_empty() {
return Ok(());
}
let adjusted: Vec<RegionCrop> = indices.iter().map(|&index| self.adjust_crop(&crops[index])).collect();
let second = self.run_pass(&adjusted, ignore)?;
for (candidate, &index) in second.into_iter().zip(indices.iter()) {
if candidate.confidence >= results[index].confidence {
results[index] = candidate;
}
}
Ok(())
}
fn adjust_crop(&self, crop: &RegionCrop) -> RegionCrop {
RegionCrop {
width: crop.width,
height: crop.height,
gray: super::contrast::adjust_contrast_grey(&crop.gray, self.config.adjust_contrast),
corners: crop.corners,
}
}
}
fn build_ignore(charset: &Charset, config: &RecognitionConfig) -> Vec<usize> {
if !config.allowlist.is_empty() {
let allowed: Vec<usize> = config.allowlist.chars().filter_map(|ch| charset.class_of(ch)).collect();
(1..charset.num_classes())
.filter(|class| !allowed.contains(class))
.collect()
} else {
config.blocklist.chars().filter_map(|ch| charset.class_of(ch)).collect()
}
}
impl TextRecognizer for CrnnRecognizer {
fn recognize(&self, crops: &[RegionCrop]) -> Result<Vec<RecognizedText>> {
if self.config.decoder != Decoder::Greedy {
return Err(OcrError::config(
"only greedy CTC decoding is implemented; set decoder = \"greedy\"",
));
}
let ignore = build_ignore(&self.charset, &self.config);
let mut results = self.run_pass(crops, &ignore)?;
self.apply_second_pass(crops, &ignore, &mut results)?;
Ok(results)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Language;
use crate::inference::Tensor;
use ndarray::{ArrayD, IxDyn};
use std::sync::atomic::{AtomicUsize, Ordering};
struct ScriptedBackend {
outputs: Vec<ArrayD<f32>>,
calls: AtomicUsize,
}
impl ScriptedBackend {
fn new(outputs: Vec<ArrayD<f32>>) -> Self {
Self {
outputs,
calls: AtomicUsize::new(0),
}
}
}
impl ModelBackend for ScriptedBackend {
fn name(&self) -> &str {
"scripted"
}
fn run(&self, _input: Tensor) -> Result<Tensor> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
let index = call.min(self.outputs.len() - 1);
Ok(self.outputs[index].clone())
}
}
fn logits(values: [f32; 3]) -> ArrayD<f32> {
ArrayD::from_shape_vec(IxDyn(&[1, 1, 3]), values.to_vec()).expect("valid logits shape")
}
fn sample_crop() -> RegionCrop {
RegionCrop {
width: 4,
height: 2,
gray: vec![120u8; 8],
corners: [[0.0, 0.0]; 4],
}
}
fn english_recognizer(backend: Arc<dyn ModelBackend>, config: RecognitionConfig) -> CrnnRecognizer {
CrnnRecognizer::new(backend, Charset::for_language(Language::English), config)
}
#[test]
fn should_decode_single_crop_to_expected_text() {
let backend = Arc::new(ScriptedBackend::new(vec![logits([0.0, 5.0, 0.0])]));
let recognizer = english_recognizer(backend, RecognitionConfig::default());
let results = recognizer.recognize(&[sample_crop()]).expect("recognition succeeds");
assert_eq!(results.len(), 1);
assert_eq!(results[0].text, "0");
assert!(results[0].confidence > 0.0, "a confident crop scores above zero");
}
#[test]
fn should_decode_multi_region_batch_in_input_order() {
let rows = [
[0.0f32, 5.0, 0.0, 0.0, 5.0, 0.0], [0.0, 0.0, 5.0, 0.0, 0.0, 5.0], [0.0, 5.0, 0.0, 0.0, 0.0, 5.0], ];
let flat: Vec<f32> = rows.iter().flatten().copied().collect();
let batch_logits = ArrayD::from_shape_vec(IxDyn(&[3, 2, 3]), flat).expect("valid batch logits shape");
let backend = Arc::new(ScriptedBackend::new(vec![batch_logits]));
let config = RecognitionConfig {
batch_size: 3,
contrast_ths: 0.0,
..RecognitionConfig::default()
};
let recognizer = english_recognizer(backend, config);
let crops = [sample_crop(), sample_crop(), sample_crop()];
let results = recognizer.recognize(&crops).expect("recognition succeeds");
let decoded: Vec<&str> = results.iter().map(|result| result.text.as_str()).collect();
assert_eq!(decoded, vec!["0", "1", "01"], "results stay in input order");
}
#[test]
fn should_replace_low_confidence_result_when_second_pass_scores_higher() {
let first = logits([(0.2f32).ln(), (0.5f32).ln(), (0.3f32).ln()]);
let second = logits([(0.05f32).ln(), (0.05f32).ln(), (0.9f32).ln()]);
let backend = Arc::new(ScriptedBackend::new(vec![first, second]));
let config = RecognitionConfig {
contrast_ths: 0.9,
..RecognitionConfig::default()
};
let recognizer = english_recognizer(backend, config);
let results = recognizer.recognize(&[sample_crop()]).expect("recognition succeeds");
assert_eq!(
results[0].text, "1",
"the higher-confidence second pass replaces the result"
);
assert!(results[0].confidence > 0.5, "confidence reflects the second pass");
}
#[test]
fn should_keep_first_result_when_second_pass_scores_lower() {
let first = logits([(0.05f32).ln(), (0.05f32).ln(), (0.9f32).ln()]);
let second = logits([(0.2f32).ln(), (0.5f32).ln(), (0.3f32).ln()]);
let backend = Arc::new(ScriptedBackend::new(vec![first, second]));
let config = RecognitionConfig {
contrast_ths: 0.95,
..RecognitionConfig::default()
};
let recognizer = english_recognizer(backend, config);
let results = recognizer.recognize(&[sample_crop()]).expect("recognition succeeds");
assert_eq!(results[0].text, "1", "the stronger first pass is retained");
}
#[test]
fn should_ignore_classes_outside_the_allowlist() {
let charset = Charset::for_language(Language::English);
let config = RecognitionConfig {
allowlist: "01".to_string(),
..RecognitionConfig::default()
};
let ignore = build_ignore(&charset, &config);
assert!(!ignore.contains(&1));
assert!(!ignore.contains(&2));
assert!(ignore.contains(&3));
}
#[test]
fn should_ignore_blocklisted_classes() {
let charset = Charset::for_language(Language::English);
let config = RecognitionConfig {
blocklist: "0".to_string(),
..RecognitionConfig::default()
};
let ignore = build_ignore(&charset, &config);
assert_eq!(ignore, vec![1]);
}
#[test]
fn should_ignore_blocklist_only_when_allowlist_is_empty() {
let charset = Charset::for_language(Language::English);
let config = RecognitionConfig {
allowlist: "01".to_string(),
blocklist: "01".to_string(),
..RecognitionConfig::default()
};
let ignore = build_ignore(&charset, &config);
assert!(!ignore.contains(&1));
assert!(!ignore.contains(&2));
}
#[test]
fn should_reject_non_greedy_decoder() {
let backend = Arc::new(ScriptedBackend::new(vec![logits([0.0, 5.0, 0.0])]));
let config = RecognitionConfig {
decoder: crate::config::Decoder::BeamSearch,
..RecognitionConfig::default()
};
let recognizer = english_recognizer(backend, config);
let result = recognizer.recognize(&[sample_crop()]);
assert!(
matches!(result, Err(OcrError::Config { .. })),
"non-greedy decoding must be rejected with a config error"
);
}
#[cfg(feature = "ort")]
#[test]
#[ignore = "requires the ONNX Runtime native library and a recognizer model file"]
fn recognize_over_real_recognizer_model() {
let model_path = std::env::var("EASYOCR_TEST_RECOG_ONNX")
.expect("set EASYOCR_TEST_RECOG_ONNX to a recognizer 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 recognizer ONNX model");
let recognizer = CrnnRecognizer::new(
Arc::from(backend),
Charset::for_language(Language::English),
RecognitionConfig::default(),
);
let crop = RegionCrop {
width: 32,
height: 16,
gray: vec![200u8; 32 * 16],
corners: [[0.0, 0.0]; 4],
};
let result = recognizer.recognize(&[crop]);
assert!(result.is_ok(), "recognition over the real model must succeed");
}
}