use crate::loader::Error; use crate::loader::error::{BuildError, ParseError, Result}; use crate::types::Session; use ddddocr_core::traits::Loader; use ort::session::Session as OrtSession; use ort::session::builder::SessionBuilder; use std::sync::{Arc, Mutex}; // pub struct OrtModelLoader; // // impl OrtModelLoader { // /// 获取针对 ORT 后端的链式构建器 // pub fn builder() -> OrtModelBuilder { // OrtModelBuilder::default() // } // } /// ORT 专用的链式构建器 #[derive(Debug, Clone)] pub struct ModelLoader { use_gpu: bool, device_id: i32, intra_threads: Option, } impl Default for ModelLoader { fn default() -> Self { Self { use_gpu: false, device_id: 0, intra_threads: None, } } } impl ModelLoader { /// 开启或关闭 GPU 加速 pub fn use_gpu(mut self, enable: bool) -> Self { self.use_gpu = enable; self } /// 指定 GPU 设备 ID pub fn device_id(mut self, id: i32) -> Self { self.device_id = id; self } pub fn num_threads(mut self, threads: usize) -> Self { self.intra_threads = Some(threads); self } /// 内部辅助方法:根据当前的配置构建 ORT 底层的 SessionBuilder fn create_session_builder(&self) -> Result { let mut builder = OrtSession::builder().map_err(|e| BuildError::BuildFailed(e))?; // 如果用户显式设置了线程数,则配置给 ORT if let Some(threads) = self.intra_threads { builder = builder .with_intra_threads(threads) .map_err(|e| BuildError::Threads(format!("设置线程数失败: {e}")))?; } if self.use_gpu { // 根据 ort 库版本配置 CUDA 执行提供者 (Execution Provider) #[cfg(feature = "cuda")] { use ort::ep::CUDAExecutionProvider; let cuda_ep = CUDAExecutionProvider::default().with_device_id(self.device_id); builder = builder .with_execution_providers([cuda_ep.build()]) .map_err(|e| { BuildError::EnabledCudaFailed(format!("配置 CUDA 硬件加速失败: {e}")) })? } #[cfg(not(feature = "cuda"))] { // 如果用户明确开启了 GPU,但 Feature 没编译进去,明确抛错提醒 return Err(BuildError::NotEnabledCuda( "未启用 CUDA 支持:请在 Cargo.toml 中为 ddddocr-ort 开启 `cuda` feature" .to_string(), ))?; } } Ok(builder) } } impl Loader for ModelLoader { type Session = Session; type Error = Error; fn build_for_path

(&self, model_path: P) -> Result where P: AsRef, { let path_ref = model_path.as_ref(); let mut builder = self.create_session_builder()?; // Session::builder() 会返回 Result let session = builder .commit_from_file(path_ref) .map_err(|e| ParseError::Path(path_ref.display().to_string(), e))?; Ok(Arc::new(Mutex::new(session))) // 这里的session需要包装下 } /// 策略 B:从内存字节流加载模型(配合 include_bytes! 使用) fn build_from_bytes(&self, model_bytes: &[u8]) -> Result { let mut builder = self.create_session_builder()?; let session = builder .commit_from_memory(model_bytes) .map_err(ParseError::Bytes)?; Ok(Arc::new(Mutex::new(session))) } }