Skip to main content

ppocr_rs/
base_net.rs

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        // [downport rc.11→rc.9] in rc.11 inputs/outputs sono metodi; in rc.9
33        // sono field pubblici di Session, e i singoli `Input` espongono
34        // `name` come campo `String` (non metodo `&str`).
35        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}