use clap::Args;
use sceptre::OcrConfig;
const MIN_PROBABILITY: f32 = 0.0;
const MAX_PROBABILITY: f32 = 1.0;
fn parse_probability(raw: &str) -> core::result::Result<f32, String> {
let value: f32 = raw.parse().map_err(|_| format!("`{raw}` is not a number"))?;
if value.is_nan() || !(MIN_PROBABILITY..=MAX_PROBABILITY).contains(&value) {
return Err(format!("must be between {MIN_PROBABILITY} and {MAX_PROBABILITY}"));
}
Ok(value)
}
#[derive(Debug, Default, Args)]
pub struct OcrOverrides {
#[arg(id = "lang", long = "lang", value_enum)]
languages: Vec<LanguageArg>,
#[arg(long)]
threads: Option<usize>,
#[arg(long, value_enum)]
backend: Option<BackendArg>,
#[arg(long, value_enum)]
accelerator: Option<AcceleratorArg>,
#[arg(long, value_parser = parse_probability)]
text_threshold: Option<f32>,
#[arg(long, value_parser = parse_probability)]
link_threshold: Option<f32>,
#[arg(long)]
canvas_size: Option<u32>,
#[arg(long)]
max_megapixels: Option<f32>,
#[arg(long)]
detect_orientation: bool,
#[arg(long)]
orientation_probe_canvas_size: Option<u32>,
#[arg(long)]
orientation_margin: Option<f32>,
}
#[derive(Debug, Clone, Copy, clap::ValueEnum)]
pub enum LanguageArg {
English,
Latin,
ChineseSimplified,
Japanese,
Korean,
Cyrillic,
Telugu,
Kannada,
}
#[derive(Debug, Clone, Copy, clap::ValueEnum)]
pub enum BackendArg {
Ort,
Tract,
Candle,
}
#[derive(Debug, Clone, Copy, clap::ValueEnum)]
pub enum AcceleratorArg {
Cpu,
Auto,
Coreml,
Directml,
Metal,
Cuda,
}
pub fn every_language() -> Vec<sceptre::Language> {
use clap::ValueEnum as _;
LanguageArg::value_variants().iter().copied().map(Into::into).collect()
}
impl From<LanguageArg> for sceptre::Language {
fn from(value: LanguageArg) -> Self {
use sceptre::Language;
match value {
LanguageArg::English => Language::English,
LanguageArg::Latin => Language::Latin,
LanguageArg::ChineseSimplified => Language::ChineseSimplified,
LanguageArg::Japanese => Language::Japanese,
LanguageArg::Korean => Language::Korean,
LanguageArg::Cyrillic => Language::Cyrillic,
LanguageArg::Telugu => Language::Telugu,
LanguageArg::Kannada => Language::Kannada,
}
}
}
impl From<BackendArg> for sceptre::Backend {
fn from(value: BackendArg) -> Self {
use sceptre::Backend;
match value {
BackendArg::Ort => Backend::Ort,
BackendArg::Tract => Backend::Tract,
BackendArg::Candle => Backend::Candle,
}
}
}
impl From<AcceleratorArg> for sceptre::Accelerator {
fn from(value: AcceleratorArg) -> Self {
use sceptre::Accelerator;
match value {
AcceleratorArg::Cpu => Accelerator::Cpu,
AcceleratorArg::Auto => Accelerator::Auto,
AcceleratorArg::Coreml => Accelerator::CoreMl,
AcceleratorArg::Directml => Accelerator::DirectMl,
AcceleratorArg::Metal => Accelerator::Metal,
AcceleratorArg::Cuda => Accelerator::Cuda,
}
}
}
impl OcrOverrides {
pub fn has_languages(&self) -> bool {
!self.languages.is_empty()
}
pub fn apply(&self, config: &mut OcrConfig) {
if !self.languages.is_empty() {
config.model.languages = self.languages.iter().copied().map(Into::into).collect();
}
if let Some(threads) = self.threads {
config.concurrency.max_threads = Some(threads);
}
if let Some(backend) = self.backend {
config.model.backend = backend.into();
}
if let Some(accelerator) = self.accelerator {
config.model.accelerator = accelerator.into();
}
if let Some(text_threshold) = self.text_threshold {
config.detection.text_threshold = text_threshold;
}
if let Some(link_threshold) = self.link_threshold {
config.detection.link_threshold = link_threshold;
}
if let Some(canvas_size) = self.canvas_size {
config.detection.canvas_size = canvas_size;
}
if let Some(max_megapixels) = self.max_megapixels {
config.detection.max_megapixels = Some(max_megapixels);
}
if self.detect_orientation {
config.detection.detect_orientation = true;
}
if let Some(orientation_probe_canvas_size) = self.orientation_probe_canvas_size {
config.detection.orientation_probe_canvas_size = orientation_probe_canvas_size;
}
if let Some(orientation_margin) = self.orientation_margin {
config.detection.orientation_margin = orientation_margin;
}
}
}
#[cfg(test)]
mod tests {
use super::{AcceleratorArg, LanguageArg, OcrOverrides, every_language, parse_probability};
#[test]
fn should_map_every_accelerator_arg_to_its_library_wire_name() {
use clap::ValueEnum as _;
let expected = ["cpu", "auto", "coreml", "directml", "metal", "cuda"];
let variants = AcceleratorArg::value_variants();
assert_eq!(variants.len(), expected.len(), "every variant must be covered");
for (variant, wire) in variants.iter().copied().zip(expected) {
assert_eq!(sceptre::Accelerator::from(variant).as_str(), wire);
}
}
#[test]
fn should_expand_every_language_to_all_value_enum_variants() {
use clap::ValueEnum as _;
let languages = every_language();
assert_eq!(languages.len(), LanguageArg::value_variants().len());
assert!(languages.contains(&sceptre::Language::English));
assert!(languages.contains(&sceptre::Language::Kannada));
let mut deduplicated = languages.clone();
deduplicated.dedup();
assert_eq!(deduplicated.len(), languages.len(), "no language may repeat");
}
#[test]
fn should_accept_probabilities_within_the_unit_interval() {
assert_eq!(parse_probability("0.0").expect("zero is valid"), 0.0);
assert_eq!(parse_probability("1.0").expect("one is valid"), 1.0);
assert_eq!(parse_probability("0.5").expect("a mid value is valid"), 0.5);
}
#[test]
fn should_reject_out_of_range_nan_and_unparseable_probabilities() {
assert!(parse_probability("-0.1").is_err(), "a negative value is rejected");
assert!(parse_probability("1.1").is_err(), "a value above one is rejected");
assert!(parse_probability("nan").is_err(), "NaN is rejected");
assert!(parse_probability("abc").is_err(), "a non-numeric value is rejected");
}
#[test]
fn should_leave_orientation_settings_at_their_default_when_unset() {
let overrides = OcrOverrides::default();
let mut config = sceptre::OcrConfig::default();
overrides.apply(&mut config);
assert!(!config.detection.detect_orientation);
assert_eq!(config.detection.orientation_probe_canvas_size, 1280);
assert_eq!(config.detection.orientation_margin, 0.05);
}
#[test]
fn should_apply_orientation_overrides_when_set() {
let overrides = OcrOverrides {
detect_orientation: true,
orientation_probe_canvas_size: Some(640),
orientation_margin: Some(0.1),
..OcrOverrides::default()
};
let mut config = sceptre::OcrConfig::default();
overrides.apply(&mut config);
assert!(config.detection.detect_orientation);
assert_eq!(config.detection.orientation_probe_canvas_size, 640);
assert_eq!(config.detection.orientation_margin, 0.1);
}
#[test]
fn should_leave_max_megapixels_unset_by_default() {
let overrides = OcrOverrides::default();
let mut config = sceptre::OcrConfig::default();
overrides.apply(&mut config);
assert_eq!(config.detection.max_megapixels, None);
}
#[test]
fn should_apply_max_megapixels_override_when_set() {
let overrides = OcrOverrides {
max_megapixels: Some(4.0),
..OcrOverrides::default()
};
let mut config = sceptre::OcrConfig::default();
overrides.apply(&mut config);
assert_eq!(config.detection.max_megapixels, Some(4.0));
}
}