use std::path::PathBuf;
use serde::{Deserialize, Serialize};
use crate::error::{OcrError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Language {
#[default]
English,
Latin,
ChineseSimplified,
Japanese,
Korean,
Cyrillic,
Telugu,
Kannada,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Backend {
#[default]
Ort,
Tract,
Candle,
}
const EVERY_BACKEND: [Backend; 3] = [Backend::Ort, Backend::Tract, Backend::Candle];
impl Backend {
pub const fn as_str(self) -> &'static str {
match self {
Self::Ort => "ort",
Self::Tract => "tract",
Self::Candle => "candle",
}
}
pub const fn hardware_accelerators(self) -> &'static [Accelerator] {
match self {
Self::Ort => &[Accelerator::CoreMl, Accelerator::DirectMl, Accelerator::Cuda],
Self::Tract => &[],
Self::Candle => &[Accelerator::Metal, Accelerator::Cuda],
}
}
pub fn supports(self, accelerator: Accelerator) -> bool {
accelerator.is_cpu_only() || self.hardware_accelerators().contains(&accelerator)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum Accelerator {
#[default]
Cpu,
Auto,
#[serde(rename = "coreml")]
CoreMl,
#[serde(rename = "directml")]
DirectMl,
Metal,
Cuda,
}
impl Accelerator {
pub const fn as_str(self) -> &'static str {
match self {
Self::Cpu => "cpu",
Self::Auto => "auto",
Self::CoreMl => "coreml",
Self::DirectMl => "directml",
Self::Metal => "metal",
Self::Cuda => "cuda",
}
}
fn equivalent_on(self, backend: Backend) -> Option<Self> {
let equivalent = match self {
Self::CoreMl => Self::Metal,
Self::Metal => Self::CoreMl,
_ => return None,
};
backend.supports(equivalent).then_some(equivalent)
}
pub(crate) const fn is_cpu_only(self) -> bool {
matches!(self, Self::Cpu | Self::Auto)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(default, deny_unknown_fields)]
pub struct ModelConfig {
pub languages: Vec<Language>,
pub backend: Backend,
pub accelerator: Accelerator,
pub detector_path: Option<PathBuf>,
pub recognizer_path: Option<PathBuf>,
pub cache_dir: Option<PathBuf>,
pub registry_owner: Option<String>,
}
impl Default for ModelConfig {
fn default() -> Self {
Self {
languages: vec![Language::English],
backend: Backend::default(),
accelerator: Accelerator::default(),
detector_path: None,
recognizer_path: None,
cache_dir: None,
registry_owner: None,
}
}
}
impl ModelConfig {
pub(crate) fn validate(&self) -> Result<()> {
match (&self.detector_path, &self.recognizer_path) {
(Some(_), None) => {
return Err(OcrError::config(
"model.recognizer_path is required when model.detector_path is configured",
));
}
(None, Some(_)) => {
return Err(OcrError::config(
"model.detector_path is required when model.recognizer_path is configured",
));
}
_ => {}
}
self.validate_accelerator()
}
fn validate_accelerator(&self) -> Result<()> {
if self.backend.supports(self.accelerator) {
return Ok(());
}
Err(OcrError::config(unsupported_accelerator(
self.backend,
self.accelerator,
)))
}
}
fn unsupported_accelerator(backend: Backend, accelerator: Accelerator) -> String {
let supported = backend.hardware_accelerators();
let mut message = if supported.is_empty() {
format!(
"model.accelerator = \"{}\" is not available: the \"{}\" backend is CPU-only",
accelerator.as_str(),
backend.as_str()
)
} else {
let names: Vec<&str> = supported.iter().map(|supported| supported.as_str()).collect();
format!(
"model.accelerator = \"{}\" is not available on the \"{}\" backend, which runs on {}",
accelerator.as_str(),
backend.as_str(),
names.join(" or ")
)
};
if let Some(equivalent) = accelerator.equivalent_on(backend) {
message.push_str(&format!(
"; the same hardware is reached with model.accelerator = \"{}\"",
equivalent.as_str()
));
} else if let Some(other) = EVERY_BACKEND.iter().find(|other| other.supports(accelerator)) {
message.push_str(&format!("; model.backend = \"{}\" runs on it", other.as_str()));
}
message
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_accept_model_paths_only_as_a_pair() {
let mut config = ModelConfig {
detector_path: Some("craft.onnx".into()),
..ModelConfig::default()
};
assert!(config.validate().is_err());
config.recognizer_path = Some("english.onnx".into());
config.validate().expect("paired paths are valid");
}
#[test]
fn should_default_the_accelerator_to_cpu() {
assert_eq!(ModelConfig::default().accelerator, Accelerator::Cpu);
assert_eq!(Accelerator::default(), Accelerator::Cpu);
}
#[test]
fn should_round_trip_every_accelerator_wire_name() {
let cases = [
(Accelerator::Cpu, "cpu"),
(Accelerator::Auto, "auto"),
(Accelerator::CoreMl, "coreml"),
(Accelerator::DirectMl, "directml"),
(Accelerator::Metal, "metal"),
(Accelerator::Cuda, "cuda"),
];
for (accelerator, wire) in cases {
let encoded = serde_json::to_string(&accelerator).expect("serialize the accelerator");
assert_eq!(
encoded,
format!("\"{wire}\""),
"unexpected wire name for {accelerator:?}"
);
assert_eq!(accelerator.as_str(), wire);
let decoded: Accelerator = serde_json::from_str(&encoded).expect("deserialize the accelerator");
assert_eq!(decoded, accelerator);
}
}
#[test]
fn should_reject_a_hardware_accelerator_on_a_cpu_only_backend() {
let config = ModelConfig {
backend: Backend::Tract,
accelerator: Accelerator::CoreMl,
..ModelConfig::default()
};
let error = config.validate().expect_err("coreml on tract must be rejected");
assert!(matches!(error, OcrError::Config { .. }), "expected a config error");
let message = error.to_string();
assert!(
message.contains("coreml"),
"message must name the accelerator: {message}"
);
assert!(message.contains("tract"), "message must name the backend: {message}");
assert!(
message.contains("CPU-only"),
"message must say why tract cannot honor it: {message}"
);
}
#[test]
fn should_reject_an_accelerator_that_belongs_to_another_backend() {
let config = ModelConfig {
backend: Backend::Candle,
accelerator: Accelerator::DirectMl,
..ModelConfig::default()
};
let error = config.validate().expect_err("directml on candle must be rejected");
let message = error.to_string();
assert!(
message.contains("directml") && message.contains("candle"),
"message must name both sides: {message}"
);
assert!(
message.contains("metal") && message.contains("cuda"),
"message must list what candle does support: {message}"
);
assert!(
message.contains("backend = \"ort\""),
"message must point at the backend that does support directml: {message}"
);
}
#[test]
fn should_name_the_apple_equivalent_when_the_wrong_framework_is_requested() {
let cases = [
(Backend::Candle, Accelerator::CoreMl, "metal"),
(Backend::Ort, Accelerator::Metal, "coreml"),
];
for (backend, accelerator, equivalent) in cases {
let config = ModelConfig {
backend,
accelerator,
..ModelConfig::default()
};
let Err(error) = config.validate() else {
panic!("{accelerator:?} on {backend:?} must be rejected");
};
let message = error.to_string();
assert!(
message.contains(equivalent),
"on the {} backend the message must point at `{equivalent}`: {message}",
backend.as_str()
);
}
}
#[test]
fn should_accept_every_accelerator_the_support_table_lists() {
for backend in [Backend::Ort, Backend::Tract, Backend::Candle] {
for accelerator in backend.hardware_accelerators() {
let config = ModelConfig {
backend,
accelerator: *accelerator,
..ModelConfig::default()
};
config
.validate()
.unwrap_or_else(|error| panic!("{accelerator:?} is listed for {backend:?} but rejected: {error}"));
assert!(backend.supports(*accelerator));
}
}
}
#[test]
fn should_list_no_hardware_accelerator_that_is_really_the_cpu() {
for backend in [Backend::Ort, Backend::Tract, Backend::Candle] {
let listed = backend.hardware_accelerators();
assert!(
listed.iter().all(|accelerator| !accelerator.is_cpu_only()),
"{backend:?} lists a CPU selection as hardware: {listed:?}"
);
assert!(
backend.supports(Accelerator::Cpu) && backend.supports(Accelerator::Auto),
"{backend:?} must accept the CPU-only selections"
);
}
}
#[test]
fn should_run_candle_on_metal_and_cuda_but_not_on_the_onnx_runtime_providers() {
assert_eq!(
Backend::Candle.hardware_accelerators(),
&[Accelerator::Metal, Accelerator::Cuda],
"candle names hardware, not ONNX Runtime execution providers"
);
assert!(!Backend::Candle.supports(Accelerator::CoreMl));
assert!(!Backend::Ort.supports(Accelerator::Metal));
assert!(Backend::Tract.hardware_accelerators().is_empty());
}
#[test]
fn should_accept_cpu_and_auto_accelerators_on_a_cpu_only_backend() {
for accelerator in [Accelerator::Cpu, Accelerator::Auto] {
let config = ModelConfig {
backend: Backend::Tract,
accelerator,
..ModelConfig::default()
};
config.validate().expect("cpu-only selections are valid on tract");
}
}
#[test]
fn should_accept_every_onnx_runtime_provider_on_the_ort_backend() {
for accelerator in [
Accelerator::Cpu,
Accelerator::Auto,
Accelerator::CoreMl,
Accelerator::DirectMl,
Accelerator::Cuda,
] {
let config = ModelConfig {
backend: Backend::Ort,
accelerator,
..ModelConfig::default()
};
config
.validate()
.expect("ort accepts every execution provider at config time");
}
}
}