use std::path::{Path, PathBuf};
use embroider::{SessionPolicy, cuda_available};
use tracing::info;
use crate::config::{ConverterConfig, RoutingMode};
use crate::error::{BobineError, Result};
use crate::rapid_layout::RapidLayout;
use crate::rapid_ocr::RapidOcr;
use crate::tex_teller::TexTeller;
pub(crate) fn session_builder(
providers: &[String],
) -> Result<ort::session::builder::SessionBuilder> {
let builder =
ort::session::Session::builder().map_err(|e| BobineError::Ort(e.to_string()))?;
let tuned = SessionPolicy::ort_defaults()
.apply(builder)
.map_err(|e| BobineError::Ort(e.to_string()))?;
Ok(embroider::apply_providers(tuned, providers))
}
const LAYOUT_REPO: (&str, &str) = ("wybxc", "DocLayout-YOLO-DocStructBench-onnx");
const LAYOUT_FILENAME: &str = "doclayout_yolo_docstructbench_imgsz1024.onnx";
const OCR_REPO: (&str, &str) = ("SWHL", "RapidOCR");
const OCR_DET_FILENAME: &str = "PP-OCRv4/ch_PP-OCRv4_det_infer.onnx";
const OCR_REC_FILENAME: &str = "PP-OCRv4/ch_PP-OCRv4_rec_infer.onnx";
fn hf_fetch(cache_dir: &Path, repo: (&str, &str), filename: &str) -> Result<PathBuf> {
let dest = cache_dir.join(filename);
if dest.exists() {
return Ok(dest);
}
info!(
repo = repo.0,
file = filename,
"Downloading model from HuggingFace..."
);
let client =
hf_hub::HFClientSync::new().map_err(|e| BobineError::Ort(format!("hf-hub init: {e}")))?;
if let Some(parent) = dest.parent() {
std::fs::create_dir_all(parent)?;
}
client
.model(repo.0, repo.1)
.download_file()
.filename(filename.to_string())
.local_dir(cache_dir.to_path_buf())
.send()
.map_err(|e| BobineError::Ort(format!("download {filename}: {e}")))?;
Ok(dest)
}
pub struct OnnxEngine {
config: ConverterConfig,
tex_teller: Option<TexTeller>,
layout: Option<RapidLayout>,
ocr: Option<RapidOcr>,
table: Option<crate::rapid_table::RapidTable>,
cache_dir: PathBuf,
layout_model_path: Option<PathBuf>,
ocr_det_path: Option<PathBuf>,
ocr_rec_path: Option<PathBuf>,
table_model_path: Option<PathBuf>,
}
impl OnnxEngine {
fn resolve_auto_gpu_providers(
&self,
override_providers: Option<&Vec<String>>,
slot: &str,
) -> Vec<String> {
match override_providers {
Some(p) => p.clone(),
None if cuda_available() => {
tracing::info!(
slot,
"CUDA registers in the loaded ONNX Runtime library: auto-enabling CUDAExecutionProvider"
);
vec![
"CUDAExecutionProvider".to_string(),
"CPUExecutionProvider".to_string(),
]
}
None => self.config.providers.ort_providers.clone(),
}
}
fn resolve_table_providers(&self) -> Vec<String> {
match self.config.providers.table_ort_providers.clone() {
Some(p) => p,
None => {
if cuda_available() {
tracing::info!(
"table slot pinned to CPUExecutionProvider (SLANet measures 2-9x slower on CUDA; set table_ort_providers to override)"
);
}
vec!["CPUExecutionProvider".to_string()]
}
}
}
pub fn new(config: &ConverterConfig, cache_dir: &Path) -> Self {
Self {
config: config.clone(),
tex_teller: None,
layout: None,
ocr: None,
table: None,
cache_dir: cache_dir.to_path_buf(),
layout_model_path: None,
ocr_det_path: None,
ocr_rec_path: None,
table_model_path: None,
}
}
pub fn set_layout_model(&mut self, path: &Path) {
self.layout_model_path = Some(path.to_path_buf());
}
pub fn set_ocr_models(&mut self, det: &Path, rec: &Path) {
self.ocr_det_path = Some(det.to_path_buf());
self.ocr_rec_path = Some(rec.to_path_buf());
}
pub fn set_table_model(&mut self, path: &Path) {
self.table_model_path = Some(path.to_path_buf());
}
pub fn ensure_models(&mut self) -> Result<()> {
if !self.config.routing.use_onnx {
return Ok(());
}
match self.config.routing.routing_mode {
RoutingMode::Never => {}
RoutingMode::Surgical => {
self.ensure_tex_teller()?;
}
_ => {
self.ensure_tex_teller()?;
self.ensure_layout()?;
self.ensure_ocr()?;
}
}
Ok(())
}
fn ensure_tex_teller(&mut self) -> Result<()> {
if self.tex_teller.is_none() {
info!(
"Loading TexTeller ({:?}) from OleehyO/TexTeller...",
self.config.models.model_precision
);
if self.config.models.model_quantization == crate::config::ModelQuantization::Int8
&& self
.config
.providers
.ort_providers
.iter()
.chain(self.config.providers.encoder_ort_providers.iter().flatten())
.chain(self.config.providers.decoder_ort_providers.iter().flatten())
.any(|p| p.to_lowercase().contains("cuda"))
{
tracing::warn!(
"model_quantization=Int8 combined with CUDA providers: quantized ops fall back across devices (Memcpy-node overhead) and measure ~2x slower than CPU. Prefer model_quantization=Fp32 when running on a GPU."
);
}
let enc_providers = self.config.providers.encoder_ort_providers.as_deref();
let dec_providers: Vec<String> = self
.config
.providers
.decoder_ort_providers
.clone()
.unwrap_or_else(|| self.config.providers.ort_providers.clone());
let tt = match self.config.models.model_quantization {
crate::config::ModelQuantization::Fp32 => TexTeller::from_pretrained_split(
"OleehyO/TexTeller",
&self.cache_dir,
self.config.models.model_precision,
enc_providers,
&dec_providers,
)?,
crate::config::ModelQuantization::Int8 => TexTeller::from_pretrained_int8_split(
&self.cache_dir,
enc_providers,
&dec_providers,
)?,
};
self.tex_teller = Some(tt);
}
Ok(())
}
fn ensure_table(&mut self) -> Result<()> {
if self.table.is_none() {
let path = match self.table_model_path.clone() {
Some(p) => p,
None => crate::rapid_table::download_slanet_plus(&self.cache_dir)?,
};
let providers = self.resolve_table_providers();
self.table = Some(crate::rapid_table::RapidTable::load(&path, &providers)?);
}
Ok(())
}
pub fn recognize_table(
&mut self,
img: &image::DynamicImage,
ocr_lines: &[crate::rapid_ocr::OcrLine],
) -> Result<Option<String>> {
self.ensure_table()?;
self.table.as_mut().unwrap().recognize(img, ocr_lines)
}
pub fn recognize_formula(&mut self, image_path: &Path) -> Result<Option<String>> {
self.recognize_formula_capped(image_path, 1024)
}
pub fn recognize_formula_capped(
&mut self,
image_path: &Path,
max_tokens: usize,
) -> Result<Option<String>> {
self.ensure_tex_teller()?;
let tt = self.tex_teller.as_mut().unwrap();
let prev = std::mem::replace(&mut tt.max_tokens, max_tokens.clamp(16, 1024));
let r = tt.recognize(image_path);
self.tex_teller.as_mut().unwrap().max_tokens = prev;
Ok(Some(r?))
}
fn ensure_layout(&mut self) -> Result<()> {
if self.layout.is_none() {
let path = match self.layout_model_path.clone() {
Some(p) => p,
None => hf_fetch(&self.cache_dir, LAYOUT_REPO, LAYOUT_FILENAME)?,
};
info!("Loading RapidLayout from {}...", path.display());
let providers = self.resolve_auto_gpu_providers(
self.config.providers.layout_ort_providers.as_ref(),
"layout",
);
let layout = RapidLayout::load(&path, &providers)?;
self.layout = Some(layout);
}
Ok(())
}
pub fn layout_regions(
&mut self,
img: &image::DynamicImage,
) -> Result<Vec<crate::rapid_layout::LayoutRegion>> {
self.ensure_layout()?;
self.layout.as_mut().unwrap().detect(img)
}
fn ensure_ocr(&mut self) -> Result<()> {
if self.ocr.is_none() {
let det = match self.ocr_det_path.clone() {
Some(p) => p,
None => hf_fetch(&self.cache_dir, OCR_REPO, OCR_DET_FILENAME)?,
};
let rec = match self.ocr_rec_path.clone() {
Some(p) => p,
None => hf_fetch(&self.cache_dir, OCR_REPO, OCR_REC_FILENAME)?,
};
info!(
"Loading RapidOCR from {} and {}...",
det.display(),
rec.display()
);
let providers = self.resolve_auto_gpu_providers(
self.config.providers.ocr_ort_providers.as_ref(),
"ocr",
);
let ocr = RapidOcr::load(&det, &rec, &self.config.models.ocr_lang, &providers)?;
self.ocr = Some(ocr);
}
Ok(())
}
pub fn ocr_lines(
&mut self,
img: &image::DynamicImage,
) -> Result<Vec<crate::rapid_ocr::OcrLine>> {
self.ensure_ocr()?;
self.ocr.as_mut().unwrap().detect_and_recognize(img)
}
pub fn convert_pdf(&mut self, path: &Path, _work_dir: &Path) -> Result<String> {
self.ensure_models()?;
let pdf_bytes = std::fs::read(path)?;
let doc = pdf_oxide::PdfDocument::from_bytes(pdf_bytes)
.map_err(|e| crate::error::BobineError::PdfOxide(format!("{:?}", e)))?;
let n_pages = doc
.page_count()
.map_err(|e| crate::error::BobineError::PdfOxide(format!("{:?}", e)))?;
let mut md_pages: Vec<String> = Vec::new();
for i in 0..n_pages {
let md = doc
.to_markdown(i as usize, &Default::default())
.unwrap_or_else(|_| doc.extract_text(i as usize).unwrap_or_default());
md_pages.push(md);
}
Ok(md_pages.join("\n\n---\n\n"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ProviderOpts;
#[test]
fn cuda_request_degrades_gracefully_without_gpu() {
if std::env::var("ORT_DYLIB_PATH").is_err() {
eprintln!("skipped: ORT_DYLIB_PATH not set");
return;
}
let result = session_builder(&[
"CUDAExecutionProvider".to_string(),
"cpu".to_string(),
]);
assert!(
result.is_ok(),
"cuda request must fall back to cpu: {}",
result
.as_ref()
.err()
.map(|e| e.to_string())
.unwrap_or_default()
);
}
#[test]
fn unknown_providers_are_skipped_and_cpu_still_works() {
if std::env::var("ORT_DYLIB_PATH").is_err() {
eprintln!("skipped: ORT_DYLIB_PATH not set");
return;
}
let result = session_builder(&[
"warp-drive".to_string(),
"CPUExecutionProvider".to_string(),
]);
assert!(result.is_ok());
}
#[test]
fn table_slot_defaults_to_cpu_even_with_cuda_base() {
let cfg = ConverterConfig {
providers: ProviderOpts {
ort_providers: vec![
"CUDAExecutionProvider".into(),
"CPUExecutionProvider".into(),
],
..Default::default()
},
..Default::default()
};
let engine = OnnxEngine::new(&cfg, Path::new("/tmp/bobine_test"));
assert_eq!(
engine.resolve_table_providers(),
vec!["CPUExecutionProvider".to_string()]
);
let cfg = ConverterConfig {
providers: ProviderOpts {
table_ort_providers: Some(vec!["CUDAExecutionProvider".into()]),
..Default::default()
},
..Default::default()
};
let engine = OnnxEngine::new(&cfg, Path::new("/tmp/bobine_test"));
assert_eq!(
engine.resolve_table_providers(),
vec!["CUDAExecutionProvider".to_string()]
);
}
#[test]
fn auto_gpu_slot_respects_pins_and_cuda_probe() {
if std::env::var("ORT_DYLIB_PATH").is_err() {
eprintln!("skipped: ORT_DYLIB_PATH not set");
return;
}
let cfg = ConverterConfig::default();
let engine = OnnxEngine::new(&cfg, Path::new("/tmp/bobine_test"));
let resolved = engine.resolve_auto_gpu_providers(None, "layout");
if cuda_available() {
assert_eq!(resolved[0], "CUDAExecutionProvider");
assert!(resolved.contains(&"CPUExecutionProvider".to_string()));
} else {
assert_eq!(resolved, vec!["CPUExecutionProvider".to_string()]);
}
let pinned = engine.resolve_auto_gpu_providers(
Some(&vec!["CPUExecutionProvider".to_string()]),
"ocr",
);
assert_eq!(pinned, vec!["CPUExecutionProvider".to_string()]);
}
}