use candle_core::{Device, Tensor as CandleTensor};
use super::craft_net::CraftNet;
use super::crnn_net::CrnnNet;
use super::onnx_proto::OnnxGraph;
use super::{candle_error, device, weights};
use crate::error::{OcrError, Result};
use crate::inference::{BackendOptions, ModelBackend, NetworkKind, Tensor, buffer};
const DETECTOR_CONVOLUTIONS: usize = 27;
const RECOGNIZER_CONVOLUTIONS: usize = 7;
const RECOGNIZER_RECURRENT_LAYERS: usize = 2;
enum Net {
Craft(Box<CraftNet>),
Crnn(Box<CrnnNet>),
}
pub(crate) struct CandleBackend {
net: Net,
device: Device,
}
impl CandleBackend {
pub(crate) fn load(model_bytes: &[u8], options: BackendOptions<'_>) -> Result<Self> {
let _ = options.threads;
let (device, accelerator) = device::select_device(options.accelerator)?;
tracing::debug!(
requested = options.accelerator.as_str(),
selected = accelerator.as_str(),
"resolved the candle device"
);
let graph = OnnxGraph::decode(model_bytes)?;
validate_network(&graph, options.network)?;
let vb = weights::var_builder(&graph, &device)?;
let net = match options.network {
NetworkKind::Detector => Net::Craft(Box::new(
CraftNet::new(vb).map_err(|error| candle_error("build the CRAFT detector", error))?,
)),
NetworkKind::Recognizer => Net::Crnn(Box::new(
CrnnNet::new(vb).map_err(|error| candle_error("build the CRNN recognizer", error))?,
)),
};
Ok(Self { net, device })
}
}
impl ModelBackend for CandleBackend {
fn name(&self) -> &str {
"candle"
}
fn run(&self, input: Tensor) -> Result<Tensor> {
let (shape, data) = buffer::input_buffer(input);
let tensor = CandleTensor::from_vec(data, shape.as_slice(), &self.device)
.map_err(|error| candle_error("build the candle input tensor", error))?;
let output = match &self.net {
Net::Craft(net) => net.forward(&tensor),
Net::Crnn(net) => net.forward(&tensor),
}
.map_err(|error| candle_error("run candle inference", error))?;
let dims = output.dims().to_vec();
let values = output
.contiguous()
.and_then(|output| output.flatten_all())
.and_then(|output| output.to_vec1::<f32>())
.map_err(|error| candle_error("read the candle output tensor as f32", error))?;
buffer::array_from_parts("candle", &dims, values)
}
}
fn validate_network(graph: &OnnxGraph, network: NetworkKind) -> Result<()> {
let (expected_convolutions, expected_recurrent) = match network {
NetworkKind::Detector => (DETECTOR_CONVOLUTIONS, 0),
NetworkKind::Recognizer => (RECOGNIZER_CONVOLUTIONS, RECOGNIZER_RECURRENT_LAYERS),
};
let convolutions = graph.op_count("Conv");
let recurrent = graph.op_count("LSTM");
if convolutions == expected_convolutions && recurrent == expected_recurrent {
return Ok(());
}
Err(OcrError::inference(format!(
"the model does not look like the {network:?} network the candle backend was asked for: \
expected {expected_convolutions} Conv and {expected_recurrent} LSTM nodes, \
found {convolutions} Conv and {recurrent} LSTM"
)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn should_reject_bytes_that_do_not_decode() {
let error = CandleBackend::load(&[0xff, 0xff], BackendOptions::default())
.err()
.expect("garbage must not load");
assert!(matches!(error, OcrError::Inference { .. }));
}
#[test]
fn should_reject_a_graph_that_does_not_match_the_requested_network() {
let graph = OnnxGraph {
nodes: Vec::new(),
initializers: std::collections::HashMap::new(),
};
let error = validate_network(&graph, NetworkKind::Recognizer).expect_err("an empty graph is not a recognizer");
let message = format!("{error}");
assert!(
message.contains("Recognizer") && message.contains("LSTM"),
"the error must name the expected role and what it looked for: {message}"
);
}
}
#[cfg(all(test, feature = "ort"))]
mod ort_parity {
use ndarray::{ArrayD, IxDyn};
use super::*;
use crate::config::{Accelerator, Backend, Language, OcrConfig};
use crate::inference::load_backend;
use crate::models::provision::model_manifest;
const TOLERANCE: f32 = 1e-4;
const LCG: (u32, u32) = (1_664_525, 1_013_904_223);
#[derive(Debug, Clone, Copy)]
enum Bar {
Absolute(f32),
#[cfg_attr(
not(feature = "candle-metal"),
expect(dead_code, reason = "constructed only by the Metal parity test")
)]
Relative(f32),
}
impl Bar {
fn allowed(self, reference: f32) -> f32 {
match self {
Self::Absolute(tolerance) => tolerance,
Self::Relative(tolerance) => tolerance * reference.abs().max(1.0),
}
}
}
fn require_models() -> bool {
std::env::var("SCEPTRE_REQUIRE_MODELS")
.map(|value| !matches!(value.trim(), "" | "0" | "false" | "no"))
.unwrap_or(false)
}
fn pseudo_random(shape: &[usize], seed: u32) -> Vec<f32> {
let count: usize = shape.iter().product();
let mut state = seed;
(0..count)
.map(|_| {
state = state.wrapping_mul(LCG.0).wrapping_add(LCG.1);
(state >> 8) as f32 / (1_u32 << 23) as f32 - 1.0
})
.collect()
}
fn cached_model(language: Language, name: &str) -> Option<std::path::PathBuf> {
let mut config = OcrConfig::default();
config.model.languages = vec![language];
let manifest = model_manifest(&config).expect("build the model manifest");
let entry = manifest.into_iter().find(|info| info.name == name)?;
entry.cached.then_some(entry.path?)
}
fn run(
backend: Backend,
accelerator: Accelerator,
bytes: &[u8],
network: NetworkKind,
shape: &[usize],
values: &[f32],
) -> (Vec<usize>, Vec<f32>) {
let options = BackendOptions {
network,
accelerator,
..BackendOptions::default()
};
let loaded = load_backend(backend, bytes, options).expect("load the model");
let input = ArrayD::from_shape_vec(IxDyn(shape), values.to_vec()).expect("build the input");
let output = loaded.run(input).expect("run inference");
(output.shape().to_vec(), output.iter().copied().collect())
}
fn assert_agreement(model: &std::path::Path, network: NetworkKind, shape: &[usize], seed: u32) {
assert_agreement_on(Accelerator::Cpu, Bar::Absolute(TOLERANCE), model, network, shape, seed);
}
fn assert_agreement_on(
accelerator: Accelerator,
bar: Bar,
model: &std::path::Path,
network: NetworkKind,
shape: &[usize],
seed: u32,
) {
let bytes = std::fs::read(model).expect("read the model file");
let values = pseudo_random(shape, seed);
let (ort_shape, ort_values) = run(Backend::Ort, Accelerator::Cpu, &bytes, network, shape, &values);
let (candle_shape, candle_values) = run(Backend::Candle, accelerator, &bytes, network, shape, &values);
assert_eq!(
ort_shape, candle_shape,
"{network:?} output shapes differ for input {shape:?}"
);
let overruns: Vec<f32> = ort_values
.iter()
.zip(candle_values.iter())
.map(|(reference, actual)| (reference - actual).abs() / bar.allowed(*reference))
.collect();
let (index, overrun) = overruns
.iter()
.copied()
.enumerate()
.max_by(|left, right| left.1.total_cmp(&right.1))
.expect("the output is non-empty");
let exceeding = overruns.iter().filter(|value| **value > 1.0).count();
let coordinates: Vec<String> = overruns
.iter()
.enumerate()
.filter(|(_, value)| **value > 1.0)
.take(8)
.map(|(flat, _)| {
let mut remainder = flat;
let mut position = Vec::new();
for extent in ort_shape.iter().rev() {
position.push(remainder % extent);
remainder /= extent;
}
position.reverse();
format!("{position:?}")
})
.collect();
assert!(
overrun <= 1.0,
"{network:?} at input {shape:?}: candle on {} differs from ort by {} at flat index \
{index} of {} (ort={}, candle={}), {:.1}x its {bar:?} bar of {}; {exceeding} values \
({:.1}%) exceed their bar, output shape {ort_shape:?}; first differing positions {coordinates:?}",
accelerator.as_str(),
(ort_values[index] - candle_values[index]).abs(),
overruns.len(),
ort_values[index],
candle_values[index],
overrun,
bar.allowed(ort_values[index]),
100.0 * exceeding as f64 / overruns.len() as f64
);
}
#[test]
fn should_agree_with_ort_on_the_craft_detector() {
let Some(model) = cached_model(Language::English, "craft_mlt_25k") else {
assert!(!require_models(), "craft_mlt_25k is not cached");
return;
};
for (index, shape) in [[1, 3, 256, 256], [1, 3, 320, 480], [1, 3, 512, 288]]
.iter()
.enumerate()
{
assert_agreement(&model, NetworkKind::Detector, shape, 1 + index as u32);
}
}
#[test]
fn should_agree_with_ort_on_the_english_recognizer() {
let Some(model) = cached_model(Language::English, "english_g2") else {
assert!(!require_models(), "english_g2 is not cached");
return;
};
for (index, shape) in [[1, 1, 64, 128], [2, 1, 64, 256], [3, 1, 64, 96]].iter().enumerate() {
assert_agreement(&model, NetworkKind::Recognizer, shape, 11 + index as u32);
}
}
#[test]
fn should_agree_with_ort_on_the_cyrillic_recognizer() {
let Some(model) = cached_model(Language::Cyrillic, "cyrillic_g2") else {
assert!(!require_models(), "cyrillic_g2 is not cached");
return;
};
assert_agreement(&model, NetworkKind::Recognizer, &[2, 1, 64, 192], 21);
}
#[cfg(feature = "candle-metal")]
#[test]
fn should_agree_with_ort_when_candle_runs_on_metal() {
let Some(detector) = cached_model(Language::English, "craft_mlt_25k") else {
assert!(!require_models(), "craft_mlt_25k is not cached");
return;
};
assert_agreement_on(
Accelerator::Metal,
Bar::Relative(TOLERANCE),
&detector,
NetworkKind::Detector,
&[1, 3, 320, 480],
2,
);
let Some(recognizer) = cached_model(Language::English, "english_g2") else {
assert!(!require_models(), "english_g2 is not cached");
return;
};
assert_agreement_on(
Accelerator::Metal,
Bar::Relative(TOLERANCE),
&recognizer,
NetworkKind::Recognizer,
&[2, 1, 64, 256],
12,
);
}
}