1use ort::session::{
2 builder::{GraphOptimizationLevel, SessionBuilder},
3 Session,
4};
5
6use crate::ocr_error::OcrError;
7
8pub trait BaseNet {
9 fn new() -> Self;
10
11 fn get_session_builder(
12 &self,
13 num_thread: usize,
14 builder_fn: Option<fn(SessionBuilder) -> Result<SessionBuilder, ort::Error>>,
15 ) -> Result<SessionBuilder, OcrError> {
16 let builder = Session::builder()?;
17 let builder = match builder_fn {
18 Some(custom) => custom(builder)?,
19 None => builder
20 .with_optimization_level(GraphOptimizationLevel::Level2)?
21 .with_intra_threads(num_thread)?
22 .with_inter_threads(num_thread)?,
23 };
24
25 Ok(builder)
26 }
27
28 fn set_input_names(&mut self, input_names: Vec<String>);
29 fn set_session(&mut self, session: Option<Session>);
30
31 fn init(&mut self, session: Session) {
32 let input_names: Vec<String> = session
36 .inputs
37 .iter()
38 .map(|input| input.name.clone())
39 .collect();
40
41 self.set_input_names(input_names);
42 self.set_session(Some(session));
43 }
44
45 fn init_model(
46 &mut self,
47 path: &str,
48 num_thread: usize,
49 builder_fn: Option<fn(SessionBuilder) -> Result<SessionBuilder, ort::Error>>,
50 ) -> Result<(), OcrError> {
51 let session = self
52 .get_session_builder(num_thread, builder_fn)?
53 .commit_from_file(path)?;
54 self.init(session);
55
56 Ok(())
57 }
58
59 fn init_model_from_memory(
60 &mut self,
61 model_bytes: &[u8],
62 num_thread: usize,
63 builder_fn: Option<fn(SessionBuilder) -> Result<SessionBuilder, ort::Error>>,
64 ) -> Result<(), OcrError> {
65 let session = self
66 .get_session_builder(num_thread, builder_fn)?
67 .commit_from_memory(model_bytes)?;
68
69 self.init(session);
70
71 Ok(())
72 }
73}