feat(model): 新增 ModelLoader 链式构建 API 及 ORT GPU/Tract 多线程配置
- 在 ddddocr-core 中定义 ModelBuilder Trait 及其错误类型 - ddddocr-ort 支持 use_gpu、device_id 及 num_threads 链式配置与 CUDA 硬件加速 - ddddocr-tract 基于 multithread-mm 特性支持 CPU 线程数控制 - 支持基于 tract-linalg 配置推理线程数,显式引入 tract-linalg 的 multithread-mm 特性,控制 GEMM 算子并发 - 优化线程池加载策略,适配 Tokio 异步及 CLI 等多场景
This commit is contained in:
@@ -2,7 +2,8 @@ use crate::types::Session;
|
||||
use ddddocr_core::error::{Result, TensorError};
|
||||
use ddddocr_core::{DetEngine, DetOutput, InferenceEngine};
|
||||
use ndarray::Ix3;
|
||||
use tract_onnx::prelude::{tvec, IntoTensor, Tensor};
|
||||
// use tract_onnx::prelude::{tvec, IntoTensor, Tensor};
|
||||
use tract_onnx::prelude::*;
|
||||
#[derive(Debug)]
|
||||
pub struct DetSession {
|
||||
pub session: Session,
|
||||
@@ -29,7 +30,8 @@ impl InferenceEngine for DetSession {
|
||||
println!("模型输出原始数据: {:?}", result);
|
||||
// Ok(result.swap_remove(0).into_tensor())
|
||||
let raw_tensor = result.swap_remove(0).into_tensor();
|
||||
let array_d = raw_tensor.into_array::<f32>().map_err(|_| {
|
||||
// raw_tensor.into_plain_array()?
|
||||
let array_d = raw_tensor.into_plain_array::<f32>().map_err(|_| {
|
||||
TensorError::Engine("Tract 实体张量无法转换为 ndarray::ArrayD".to_string())
|
||||
})?;
|
||||
// 提前利用克隆(Clone)备份好当前未转维度前的真实 shape (Vec<usize>)
|
||||
|
||||
@@ -4,7 +4,6 @@ mod ocr;
|
||||
mod types;
|
||||
mod error;
|
||||
|
||||
pub use ddddocr_core::ocr::OcrBuilder;
|
||||
pub use ddddocr_core::{SlideResult, Slider};
|
||||
pub use ddddocr_core::{SlideResult, Slider,OcrBuilder};
|
||||
pub use det::session::DetSession;
|
||||
pub use ocr::session::OcrSession;
|
||||
|
||||
@@ -3,5 +3,5 @@ mod metadata;
|
||||
mod model;
|
||||
|
||||
pub use error::{Error, ParseError, Result};
|
||||
pub use metadata::ModelMetadata;
|
||||
pub use model::ModelLoader;
|
||||
pub use metadata::{ModelMetadataDto, NormalizationDto, TractModelMetadata};
|
||||
pub use model::TractModelLoader;
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
use crate::loader::error::{Error, Result};
|
||||
|
||||
pub use ddddocr_core::ModelMetadata;
|
||||
use ddddocr_core::ocr::Resize;
|
||||
use ddddocr_core::ModelMetadata;
|
||||
use ddddocr_core::Resize;
|
||||
use ddddocr_core::{Charset, Normalization};
|
||||
use serde::Deserialize;
|
||||
use std::borrow::Cow;
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(rename_all = "snake_case")] // 支持 json 中写 "zero_to_one" 或 "minus_one_to_one"
|
||||
enum NormalizationDto {
|
||||
pub enum NormalizationDto {
|
||||
/// 映射到 [0.0, 1.0] -> pixel / 255.0
|
||||
ZeroToOne,
|
||||
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
|
||||
@@ -26,7 +26,7 @@ impl From<NormalizationDto> for Normalization {
|
||||
|
||||
/// 仅用于反序列化 JSON 的中间临时结构体(DTO)
|
||||
#[derive(Deserialize)]
|
||||
struct ModelMetadataDto {
|
||||
pub struct ModelMetadataDto {
|
||||
charset: Vec<String>,
|
||||
word: bool,
|
||||
#[serde(alias = "image")]
|
||||
|
||||
@@ -1,17 +1,54 @@
|
||||
use crate::loader::error;
|
||||
use crate::loader::error::{Error, ParseError, Result};
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::ModelBuilder;
|
||||
use std::io::Cursor;
|
||||
use tract_linalg::multithread::{Executor, set_default_executor};
|
||||
use tract_onnx::onnx;
|
||||
use tract_onnx::prelude::*;
|
||||
|
||||
pub struct ModelLoader;
|
||||
pub struct TractModelLoader;
|
||||
impl TractModelLoader {
|
||||
/// 获取针对 Tract 后端的链式构建器
|
||||
pub fn builder() -> TractModelBuilder {
|
||||
TractModelBuilder::default()
|
||||
}
|
||||
}
|
||||
|
||||
impl ModelLoader {
|
||||
pub fn model_for_path<P>(model_path: P) -> Result<Session>
|
||||
/// Tract 专用的链式构建器
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct TractModelBuilder {
|
||||
num_threads: Option<usize>,
|
||||
}
|
||||
|
||||
impl TractModelBuilder {
|
||||
/// 可选扩展:设置 CPU 线程数(不提供任何 GPU 相关的 API)
|
||||
pub fn num_threads(mut self, threads: usize) -> Self {
|
||||
self.num_threads = Some(threads);
|
||||
self
|
||||
}
|
||||
fn setup_tract_threads(&self) {
|
||||
// 💡 1. 如果设置了线程数,可以通过 tract 的 multithread 配置应用给 model
|
||||
if let Some(threads) = self.num_threads {
|
||||
// 在 Tract 中可以通过 set_num_threads 或设置底层环境控制并发
|
||||
// (注:Tract 0.20+ 版本支持全局/局部线程控制)
|
||||
let executor = if threads <= 1 {
|
||||
Executor::SingleThread
|
||||
} else {
|
||||
Executor::multithread(threads)
|
||||
};
|
||||
set_default_executor(executor);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ModelBuilder for TractModelBuilder {
|
||||
type Session = Session;
|
||||
type Error = Error;
|
||||
fn model_for_path<P>(&self, model_path: P) -> Result<Session>
|
||||
where
|
||||
P: AsRef<std::path::Path>,
|
||||
{
|
||||
self.setup_tract_threads();
|
||||
let path_ref = model_path.as_ref();
|
||||
|
||||
let session = onnx()
|
||||
@@ -27,7 +64,8 @@ impl ModelLoader {
|
||||
Ok(session)
|
||||
}
|
||||
/// 策略 B:从内存字节流加载模型(配合 include_bytes! 使用)
|
||||
pub fn model_from_bytes(model_bytes: &[u8]) -> Result<Session> {
|
||||
fn model_from_bytes(&self, model_bytes: &[u8]) -> Result<Session> {
|
||||
self.setup_tract_threads();
|
||||
// 使用 std::io::Cursor 将 &[u8] 包装为可读的流(实现 std::io::Read)
|
||||
let mut cursor = Cursor::new(model_bytes);
|
||||
|
||||
@@ -52,7 +90,7 @@ mod tests {
|
||||
|
||||
/// 辅助函数:动态构建一个简单的 ONNX/Tract 内存模型图用于测试
|
||||
fn create_test_model() -> std::result::Result<Session, anyhow::Error> {
|
||||
let mut rect = tract_onnx::prelude::Graph::default();
|
||||
let mut rect = TypedModel::default();
|
||||
|
||||
// 0.21.10 最稳妥的静态 Fact 构建
|
||||
let input_fact = TypedFact::dt_shape(DatumType::F32, &[1, 3, 224, 224]);
|
||||
@@ -63,8 +101,6 @@ mod tests {
|
||||
|
||||
rect.set_input_outlets(&[input_node.into()])
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
rect.set_output_outlets(&[input_node.into()])
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
|
||||
let typed = rect
|
||||
.into_optimized()
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use crate::loader::ModelMetadata;
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::ModelMetadata;
|
||||
use ddddocr_core::error::{DdddError, Result, TensorError};
|
||||
use ddddocr_core::utils::normalize_ocr_logits;
|
||||
use ddddocr_core::{InferenceEngine, OcrEngine, OcrOutput};
|
||||
@@ -83,7 +83,7 @@ impl InferenceEngine for OcrSession {
|
||||
match raw_tensor.datum_type() {
|
||||
DatumType::I64 => {
|
||||
let array_d = raw_tensor
|
||||
.into_array::<i64>()
|
||||
.into_plain_array::<i64>()
|
||||
.map_err(|_| TensorError::Engine("Tract 无法获取 i64 内存视图".to_string()))?;
|
||||
// .context("Tract 无法获取 i64 内存视图")?;
|
||||
// 🌟 提前提取真实维度
|
||||
@@ -101,8 +101,9 @@ impl InferenceEngine for OcrSession {
|
||||
DatumType::F32 => {
|
||||
let shape = raw_tensor.shape();
|
||||
println!("模型输出shape数据: {:?}", shape);
|
||||
// raw_tensor.to_plain_array_view()
|
||||
let view = raw_tensor
|
||||
.to_array_view::<f32>()
|
||||
.to_plain_array_view::<f32>()
|
||||
.map_err(|_| TensorError::Engine("Tract 无法获取 f32 内存视图".to_string()))?;
|
||||
// 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗
|
||||
normalize_ocr_logits(view, shape)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use tract_onnx::prelude::{Graph, RunnableModel, TypedFact, TypedOp};
|
||||
use std::sync::Arc;
|
||||
use tract_onnx::prelude::TypedRunnableModel;
|
||||
|
||||
pub type Session = RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>;
|
||||
pub type Session = Arc<TypedRunnableModel>;
|
||||
|
||||
Reference in New Issue
Block a user