refactor(load,tract):将 ModelMetadata JSON 加载逻辑解耦至 ddddocr-tract, 优化 Error 枚举结构与错误透传
- 在 load 模块中精简 Error 与 Result 别名定义 - 增加 ParseError 子类型区分路径与字节流加载失败 - 支持通过 #[from] 自动转换 Tract 引擎底层错误 - 移出 core 中的 serde 依赖,保持核心库纯洁 - 在 tract 中实现 TractModelMetadata 扩展 trait 加载解析配置
This commit is contained in:
@@ -1,68 +1,45 @@
|
||||
use crate::loader::ModelLoader;
|
||||
use anyhow::Context;
|
||||
use ddddocr_core::error::{DdddError, Result, TensorErrorReason};
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::error::{Result, TensorError};
|
||||
use ddddocr_core::{DetEngine, DetOutput, InferenceEngine};
|
||||
use ndarray::Ix3;
|
||||
use std::path::Path;
|
||||
use tract_onnx::prelude::{Graph, IntoTensor, RunnableModel, Tensor, TypedFact, TypedOp, tvec};
|
||||
use tract_onnx::prelude::{tvec, IntoTensor, Tensor};
|
||||
#[derive(Debug)]
|
||||
pub struct DetSession {
|
||||
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
|
||||
pub session: Session,
|
||||
}
|
||||
|
||||
impl DetSession {
|
||||
pub fn new<P>(model_path: P) -> Result<Self>
|
||||
where
|
||||
P: AsRef<Path>,
|
||||
{
|
||||
let session = ModelLoader::model_for_path(&model_path)?.session;
|
||||
Ok(Self { session })
|
||||
pub fn new(session: Session) -> Self {
|
||||
Self { session }
|
||||
}
|
||||
|
||||
pub fn model_from_bytes(model_bytes: &[u8]) -> Result<Self> {
|
||||
let session = ModelLoader::model_from_bytes(model_bytes)?.session;
|
||||
Ok(Self { session })
|
||||
}
|
||||
// pub fn inference(&self, tensor: Tensor) -> anyhow::Result<Tensor> {
|
||||
// // tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
|
||||
// // let result = self.ocr.run(tvec!(tensor.into()))?;
|
||||
// let mut result = self
|
||||
// .session
|
||||
// .run(tvec!(tensor.into()))
|
||||
// .context("执行模型推理失败")?;
|
||||
// println!("模型输出原始数据: {:?}", result);
|
||||
// Ok(result.swap_remove(0).into_tensor())
|
||||
// }
|
||||
}
|
||||
|
||||
impl InferenceEngine for DetSession {
|
||||
type Output = DetOutput; // 明确绑定 OCR 小枚举
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output> {
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
|
||||
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
|
||||
// let result = self.ocr.run(tvec!(tensor.into()))?;
|
||||
let tensor = Tensor::from(input_array);
|
||||
|
||||
let mut result = self.session.run(tvec!(tensor.into())).map_err(|_| {
|
||||
DdddError::Inference(TensorErrorReason::EngineError(
|
||||
"执行模型推理失败".to_string(),
|
||||
))
|
||||
})?;
|
||||
let mut result = self
|
||||
.session
|
||||
.run(tvec!(tensor.into()))
|
||||
.map_err(|_| TensorError::Engine("执行模型推理失败".to_string()))?;
|
||||
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(|_| {
|
||||
DdddError::Inference(TensorErrorReason::EngineError(
|
||||
"Tract 实体张量无法转换为 ndarray::ArrayD".to_string(),
|
||||
))
|
||||
TensorError::Engine("Tract 实体张量无法转换为 ndarray::ArrayD".to_string())
|
||||
})?;
|
||||
// 提前利用克隆(Clone)备份好当前未转维度前的真实 shape (Vec<usize>)
|
||||
let actual_shape = array_d.shape().to_vec();
|
||||
|
||||
let array3 = array_d.into_dimensionality::<Ix3>().map_err(|_| {
|
||||
DdddError::Inference(TensorErrorReason::TensorDimensionMismatch {
|
||||
TensorError::DimensionMismatch {
|
||||
expected: "3D 检测矩阵 [Batch, Box_Count, Box_Attributes]".to_string(),
|
||||
actual: actual_shape, // 优雅降维失败时动态捕获
|
||||
})
|
||||
}
|
||||
})?;
|
||||
Ok(DetOutput::Detection(array3))
|
||||
|
||||
|
||||
8
ddddocr-tract/src/error.rs
Normal file
8
ddddocr-tract/src/error.rs
Normal file
@@ -0,0 +1,8 @@
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Error, Debug)]
|
||||
pub enum TensorError {
|
||||
/// 替换原有的 anyhow::Error,明确将 Tract/ONNX 引擎底层报错序列化为干净的 String
|
||||
#[error("推理引擎内部发生异常: {0}")]
|
||||
Engine(String),
|
||||
}
|
||||
@@ -1,6 +1,10 @@
|
||||
mod det;
|
||||
pub mod loader;
|
||||
mod ocr;
|
||||
mod types;
|
||||
mod error;
|
||||
|
||||
pub use ddddocr_core::ocr::OcrBuilder;
|
||||
pub use ddddocr_core::{SlideResult, Slider};
|
||||
pub use det::session::DetSession;
|
||||
pub use ocr::session::OcrSession;
|
||||
pub use ocr::session::OcrSession;
|
||||
|
||||
@@ -1,214 +1,7 @@
|
||||
use anyhow::Context;
|
||||
use ddddocr_core::error::{DdddError, Result};
|
||||
use std::fmt;
|
||||
use std::io::Cursor;
|
||||
use tract_onnx::onnx;
|
||||
use tract_onnx::prelude::*;
|
||||
// 引入核心层的统一错误类型
|
||||
/// 明确命名为 AxisDim,代表模型某一个轴的维度特征
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub enum AxisDim {
|
||||
/// 静态固定维度(如通道数固定为 1,高度固定为 64)
|
||||
Static(usize),
|
||||
/// 动态符号维度(如宽度是动态的 "image_width")
|
||||
Dynamic(String),
|
||||
}
|
||||
mod error;
|
||||
mod metadata;
|
||||
mod model;
|
||||
|
||||
impl AxisDim {
|
||||
/// 便捷方法:判断是否为动态维度
|
||||
pub fn is_dynamic(&self) -> bool {
|
||||
matches!(self, AxisDim::Dynamic(_))
|
||||
}
|
||||
}
|
||||
/// 自定义 Debug 格式化输出,彻底融化套娃外壳,保证日志干净漂亮
|
||||
impl fmt::Debug for AxisDim {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
AxisDim::Static(size) => write!(f, "{}", size),
|
||||
AxisDim::Dynamic(expr) => write!(f, "Dynamic(\"{}\")", expr),
|
||||
}
|
||||
}
|
||||
}
|
||||
/// 模拟 Python 的 input_info 和 output_info 结构
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TensorInfo {
|
||||
pub name: String,
|
||||
pub shape: Vec<AxisDim>, // 既包含 Fixed 静态维度,也包含 Dynamic 动态符号
|
||||
pub data_type: DatumType, // 对应 Python 的 type
|
||||
}
|
||||
|
||||
/// 最终返回的模型完整信息
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ModelInfo {
|
||||
pub inputs: Vec<TensorInfo>,
|
||||
pub outputs: Vec<TensorInfo>,
|
||||
/// 硬件执行提供者(采用 Option 兼容不同底层的推理引擎)
|
||||
pub providers: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
pub struct ModelLoader {
|
||||
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
|
||||
}
|
||||
|
||||
impl ModelLoader {
|
||||
pub fn model_for_path<P>(model_path: P) -> Result<Self>
|
||||
where
|
||||
P: AsRef<std::path::Path>,
|
||||
{
|
||||
let session = onnx()
|
||||
.model_for_path(model_path).map_err(DdddError::new)?
|
||||
// .with_context(|| "加载 ONNX 模型失败,请检查路径是否正确")?
|
||||
.into_optimized().map_err(DdddError::new)?
|
||||
// .with_context(|| "优化 Tract 模型图失败")?
|
||||
.into_runnable().map_err(DdddError::new)?;
|
||||
// .with_context(|| "构建可运行 Tract 实例失败")?;
|
||||
Ok(Self { session })
|
||||
}
|
||||
/// 策略 B:从内存字节流加载模型(配合 include_bytes! 使用)
|
||||
pub fn model_from_bytes(model_bytes: &[u8]) -> Result<Self> {
|
||||
// 使用 std::io::Cursor 将 &[u8] 包装为可读的流(实现 std::io::Read)
|
||||
let mut cursor = Cursor::new(model_bytes);
|
||||
|
||||
let session = onnx()
|
||||
.model_for_read(&mut cursor)
|
||||
.map_err(DdddError::new)?
|
||||
// .with_context(|| "从内存字节流解析 ONNX 模型失败")?
|
||||
.into_optimized()
|
||||
.map_err(DdddError::new)?
|
||||
// .with_context(|| "优化 Tract 模型图失败")?
|
||||
.into_runnable().map_err(DdddError::new)?;
|
||||
// .with_context(|| "构建可运行 Tract 实例失败")?;
|
||||
|
||||
Ok(Self { session })
|
||||
}
|
||||
}
|
||||
impl ModelLoader {
|
||||
/// 获取模型详细元数据信息(对标 Python ddddocr 的 get_model_info)
|
||||
/// 完美包容 [1, 1, 64, image_width] 这样的变长图像模型
|
||||
/// 获取模型详细元数据信息(代码更紧凑、优雅)
|
||||
pub fn model_info(&self) -> Result<ModelInfo> {
|
||||
let model = self.session.model();
|
||||
|
||||
// 使用私有辅助函数统一处理,消除重复代码
|
||||
let inputs = self.resolve_tensors(
|
||||
model
|
||||
.input_outlets().map_err(DdddError::new)?
|
||||
// .map_err(|e| DdddError::InternalError(format!("获取输入节点失败: {:?}", e)))?,
|
||||
)?;
|
||||
let outputs = self.resolve_tensors(
|
||||
model
|
||||
.output_outlets().map_err(DdddError::new)?
|
||||
// .map_err(|e| DdddError::InternalError(format!("获取输出节点失败: {:?}", e)))?,
|
||||
)?;
|
||||
|
||||
Ok(ModelInfo {
|
||||
inputs,
|
||||
outputs,
|
||||
providers: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// 提取出来的公共转换逻辑:将一组 OutletId 解析为 TensorInfo 列表
|
||||
fn resolve_tensors(&self, outlets: &[OutletId]) -> Result<Vec<TensorInfo>> {
|
||||
let model = self.session.model();
|
||||
|
||||
outlets
|
||||
.iter()
|
||||
.map(|&outlet_id| {
|
||||
let fact = model.outlet_fact(outlet_id).map_err(DdddError::new)?;
|
||||
// .map_err(|e| {
|
||||
// DdddError::InternalError(format!("解析节点 Fact 失败: {:?}", e))
|
||||
// })?;
|
||||
|
||||
let shape = self.resolve_shape(&fact.shape)?;
|
||||
let node_name = model.node(outlet_id.node).name.clone();
|
||||
|
||||
Ok(TensorInfo {
|
||||
name: node_name,
|
||||
shape,
|
||||
data_type: fact.datum_type,
|
||||
})
|
||||
})
|
||||
.collect() // 函数式声明:自动传播第一处发生的错误
|
||||
}
|
||||
|
||||
/// 安全还原 Tract 维度至 Vec<AxisDim>
|
||||
fn resolve_shape(&self, shape_fact: &ShapeFact) -> Result<Vec<AxisDim>> {
|
||||
let tract_shape = shape_fact.to_tvec();
|
||||
|
||||
let resolved = tract_shape
|
||||
.iter()
|
||||
.map(|dim| {
|
||||
// 防御性编程:必须同时满足能够转换为 i64 且 大于等于 0
|
||||
if let Ok(size) = dim.to_i64() {
|
||||
if size >= 0 {
|
||||
AxisDim::Static(size as usize)
|
||||
} else {
|
||||
// 如果 ONNX 导出时某些动态维度被标记为了 -1,安全地作为动态符号捕获
|
||||
AxisDim::Dynamic(dim.to_string())
|
||||
}
|
||||
} else {
|
||||
AxisDim::Dynamic(dim.to_string())
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(resolved)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// 辅助函数:动态构建一个简单的 ONNX/Tract 内存模型图用于测试
|
||||
fn create_test_model() -> std::result::Result<
|
||||
RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
|
||||
anyhow::Error,
|
||||
> {
|
||||
let mut rect = tract_onnx::prelude::Graph::default();
|
||||
|
||||
// 0.21.10 最稳妥的静态 Fact 构建
|
||||
let input_fact = TypedFact::dt_shape(DatumType::F32, &[1, 3, 224, 224]);
|
||||
|
||||
let input_node = rect
|
||||
.add_source("input", input_fact)
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
|
||||
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()
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
let runnable = typed
|
||||
.into_runnable()
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
Ok(runnable)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_input_output_shapes_and_type() {
|
||||
let session = create_test_model().expect("建立测试模型图失败");
|
||||
let loader = ModelLoader { session };
|
||||
println!("{:?}", loader.model_info().unwrap());
|
||||
// 1. 测试输入维度解析
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_resolve_shape_logic_direct() {
|
||||
// 创建一个哑 ModelLoader 实例(session 用不上,因为我们直接测私有方法)
|
||||
let session = create_test_model().expect("建立测试模型图失败");
|
||||
let loader = ModelLoader { session };
|
||||
let dims: Vec<TDim> = vec![TDim::from(1), TDim::from(3), TDim::from(224)];
|
||||
// 方案二的精髓:我们直接利用已导出的 ShapeFact 来纯手工验证边界逻辑!
|
||||
// 1. 验证纯静态维度是否被正确还原
|
||||
let static_shape = ShapeFact::from_dims(dims);
|
||||
|
||||
let res = loader
|
||||
.resolve_shape(&static_shape)
|
||||
.expect("解析静态 shape 失败");
|
||||
}
|
||||
}
|
||||
pub use error::{Error, ParseError, Result};
|
||||
pub use metadata::ModelMetadata;
|
||||
pub use model::ModelLoader;
|
||||
|
||||
50
ddddocr-tract/src/loader/error.rs
Normal file
50
ddddocr-tract/src/loader/error.rs
Normal file
@@ -0,0 +1,50 @@
|
||||
use tract_onnx::prelude::TractError;
|
||||
pub type Result<T> = std::result::Result<T, Error>;
|
||||
#[derive(thiserror::Error, Debug)]
|
||||
pub enum Error {
|
||||
/// 解析 ONNX 模型/路径失败(如文件损坏、算子不支持、路径非法)
|
||||
#[error("解析 ONNX 模型结构失败: {0}")]
|
||||
ModelParse(#[from] ParseError),
|
||||
|
||||
/// 模型计算图优化失败(如常量折叠、形状推导失败)
|
||||
#[error("优化 Tract 模型图失败: {0}")]
|
||||
OptimizationFailed(#[source] TractError),
|
||||
|
||||
/// 构建可执行 Session 失败(如输入输出 Tensor 类型/形状未确定)
|
||||
#[error("构建可运行 Tract 实例失败: {0}")]
|
||||
RunnableBuildFailed(#[source] TractError),
|
||||
|
||||
/// JSON 反序列化失败(自动透传 serde_json 报错)
|
||||
#[error("模型 Metadata JSON 解析失败: {0}")]
|
||||
JsonParse(#[from] serde_json::Error),
|
||||
|
||||
/// 字节流非合法 UTF-8 编码(自动透传 Utf8Error)
|
||||
#[error("Metadata 字节流不是合法的 UTF-8 编码: {0}")]
|
||||
InvalidUtf8(#[from] std::str::Utf8Error),
|
||||
|
||||
#[error("模型元数据解析失败: {0}")]
|
||||
MetadataParse(String),
|
||||
|
||||
/// 承载任何第三方扩展、解密、特定预处理插件在执行时产生的自定义错误
|
||||
#[error("{0}: {1}")]
|
||||
Other(String, #[source] Box<dyn std::error::Error + Send + Sync>),
|
||||
}
|
||||
impl Error {
|
||||
/// 方便将任何第三方 Error 包装为 Error::Other
|
||||
pub fn new<E>(msg: impl Into<String>, err: E) -> Self
|
||||
where
|
||||
E: Into<Box<dyn std::error::Error + Send + Sync>>,
|
||||
{
|
||||
Self::Other(msg.into(), err.into())
|
||||
}
|
||||
}
|
||||
#[derive(thiserror::Error,Debug)]
|
||||
pub enum ParseError{
|
||||
/// 策略 A:从文件路径加载失败(附带路径上下文信息,方便排查是找不到文件还是格式不对)
|
||||
#[error("从路径 '{0}' 加载 ONNX 模型失败: {1}")]
|
||||
Path(String, #[source] TractError),
|
||||
|
||||
/// 策略 B:从内存字节流加载失败(如 include_bytes! 传入的字节流损坏)
|
||||
#[error("从内存字节流解析 ONNX 模型失败: {0}")]
|
||||
Bytes(#[source] TractError),
|
||||
}
|
||||
93
ddddocr-tract/src/loader/metadata.rs
Normal file
93
ddddocr-tract/src/loader/metadata.rs
Normal file
@@ -0,0 +1,93 @@
|
||||
use crate::loader::error::{Error, Result};
|
||||
|
||||
pub use ddddocr_core::ModelMetadata;
|
||||
use ddddocr_core::ocr::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 {
|
||||
/// 映射到 [0.0, 1.0] -> pixel / 255.0
|
||||
ZeroToOne,
|
||||
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
|
||||
MinusOneToOne,
|
||||
}
|
||||
|
||||
impl From<NormalizationDto> for Normalization {
|
||||
fn from(dto: NormalizationDto) -> Self {
|
||||
match dto {
|
||||
NormalizationDto::ZeroToOne => Normalization::ZeroToOne,
|
||||
NormalizationDto::MinusOneToOne => Normalization::MinusOneToOne,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 仅用于反序列化 JSON 的中间临时结构体(DTO)
|
||||
#[derive(Deserialize)]
|
||||
struct ModelMetadataDto {
|
||||
charset: Vec<String>,
|
||||
word: bool,
|
||||
#[serde(alias = "image")]
|
||||
resize: Vec<i32>,
|
||||
channel: u8,
|
||||
/// 新增:允许在配置文件中指定归一化策略。
|
||||
/// 使用 serde(default) 可以在不配置时提供一个默认值(比如默认 ZeroToOne)
|
||||
#[serde(default = "default_normalization")]
|
||||
normalization: NormalizationDto,
|
||||
}
|
||||
fn default_normalization() -> NormalizationDto {
|
||||
NormalizationDto::ZeroToOne
|
||||
}
|
||||
|
||||
/// Tract 专属扩展trait 或 工具函数
|
||||
pub trait TractModelMetadata: Sized {
|
||||
fn from_json_str(json_str: &str) -> Result<Self>;
|
||||
/// 机制 2:从内存字节流加载(极大地方便 include_bytes! 或网络下载)
|
||||
fn from_json_bytes(bytes: &[u8]) -> Result<Self> {
|
||||
let json_str = std::str::from_utf8(bytes)?;
|
||||
Self::from_json_str(json_str)
|
||||
}
|
||||
}
|
||||
impl TractModelMetadata for ModelMetadata {
|
||||
// --- 优雅的工厂模式构造器 ---
|
||||
fn from_json_str(json_str: &str) -> Result<ModelMetadata> {
|
||||
let dto: ModelMetadataDto = serde_json::from_str(json_str)?;
|
||||
|
||||
// 1. 将 DTO 的字符串数组转化为强类型的 Charset
|
||||
let tokens: Vec<Cow<'static, str>> =
|
||||
dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
|
||||
let charset = Charset::new(tokens);
|
||||
|
||||
// 2. 解析 resize 策略(重现 Python 的复杂条件判断)
|
||||
if dto.resize.len() != 2 {
|
||||
return Err(Error::MetadataParse(
|
||||
"'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]".to_string(),
|
||||
));
|
||||
}
|
||||
let r0 = dto.resize[0];
|
||||
let r1 = dto.resize[1];
|
||||
|
||||
let resize = if r0 == -1 {
|
||||
if dto.word {
|
||||
// 如果 word 为 true,且包含 -1,Python 里是 resize 为 (r1, r1) 的正方形
|
||||
Resize::Square(r1 as u32)
|
||||
} else {
|
||||
// 如果 word 为 false,且包含 -1,Python 里是高度固定为 r1,宽度按原图比例缩放
|
||||
Resize::DynamicWidth(r1 as u32)
|
||||
}
|
||||
} else {
|
||||
// 正常的固定宽高
|
||||
Resize::Fixed(r0 as u32, r1 as u32)
|
||||
};
|
||||
|
||||
Ok(ModelMetadata::new(
|
||||
charset,
|
||||
dto.word,
|
||||
resize,
|
||||
dto.channel,
|
||||
dto.normalization.into(),
|
||||
))
|
||||
}
|
||||
}
|
||||
99
ddddocr-tract/src/loader/model.rs
Normal file
99
ddddocr-tract/src/loader/model.rs
Normal file
@@ -0,0 +1,99 @@
|
||||
use crate::loader::error;
|
||||
use crate::loader::error::{Error, ParseError, Result};
|
||||
use crate::types::Session;
|
||||
use std::io::Cursor;
|
||||
use tract_onnx::onnx;
|
||||
use tract_onnx::prelude::*;
|
||||
|
||||
pub struct ModelLoader;
|
||||
|
||||
impl ModelLoader {
|
||||
pub fn model_for_path<P>(model_path: P) -> Result<Session>
|
||||
where
|
||||
P: AsRef<std::path::Path>,
|
||||
{
|
||||
let path_ref = model_path.as_ref();
|
||||
|
||||
let session = onnx()
|
||||
.model_for_path(path_ref)
|
||||
.map_err(|e| ParseError::Path(path_ref.display().to_string(), e))?
|
||||
// .with_context(|| "加载 ONNX 模型失败,请检查路径是否正确")?
|
||||
.into_optimized()
|
||||
.map_err(Error::OptimizationFailed)?
|
||||
// .with_context(|| "优化 Tract 模型图失败")?
|
||||
.into_runnable()
|
||||
.map_err(Error::RunnableBuildFailed)?;
|
||||
// .with_context(|| "构建可运行 Tract 实例失败")?;
|
||||
Ok(session)
|
||||
}
|
||||
/// 策略 B:从内存字节流加载模型(配合 include_bytes! 使用)
|
||||
pub fn model_from_bytes(model_bytes: &[u8]) -> Result<Session> {
|
||||
// 使用 std::io::Cursor 将 &[u8] 包装为可读的流(实现 std::io::Read)
|
||||
let mut cursor = Cursor::new(model_bytes);
|
||||
|
||||
let session = onnx()
|
||||
.model_for_read(&mut cursor)
|
||||
.map_err(ParseError::Bytes)?
|
||||
// .with_context(|| "从内存字节流解析 ONNX 模型失败")?
|
||||
.into_optimized()
|
||||
.map_err(Error::OptimizationFailed)?
|
||||
// .with_context(|| "优化 Tract 模型图失败")?
|
||||
.into_runnable()
|
||||
.map_err(Error::RunnableBuildFailed)?;
|
||||
// .with_context(|| "构建可运行 Tract 实例失败")?;
|
||||
|
||||
Ok(session)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
/// 辅助函数:动态构建一个简单的 ONNX/Tract 内存模型图用于测试
|
||||
fn create_test_model() -> std::result::Result<Session, anyhow::Error> {
|
||||
let mut rect = tract_onnx::prelude::Graph::default();
|
||||
|
||||
// 0.21.10 最稳妥的静态 Fact 构建
|
||||
let input_fact = TypedFact::dt_shape(DatumType::F32, &[1, 3, 224, 224]);
|
||||
|
||||
let input_node = rect
|
||||
.add_source("input", input_fact)
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
|
||||
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()
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
let runnable = typed
|
||||
.into_runnable()
|
||||
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
|
||||
Ok(runnable)
|
||||
}
|
||||
|
||||
// #[test]
|
||||
// fn test_input_output_shapes_and_type() {
|
||||
// let session = create_test_model().expect("建立测试模型图失败");
|
||||
//
|
||||
// println!("{:?}", ModelLoader::model_info(&session).unwrap());
|
||||
// // 1. 测试输入维度解析
|
||||
// }
|
||||
//
|
||||
// #[test]
|
||||
// fn test_resolve_shape_logic_direct() {
|
||||
// // 创建一个哑 ModelLoader 实例(session 用不上,因为我们直接测私有方法)
|
||||
// let session = create_test_model().expect("建立测试模型图失败");
|
||||
//
|
||||
// let dims: Vec<TDim> = vec![TDim::from(1), TDim::from(3), TDim::from(224)];
|
||||
// // 方案二的精髓:我们直接利用已导出的 ShapeFact 来纯手工验证边界逻辑!
|
||||
// // 1. 验证纯静态维度是否被正确还原
|
||||
// let static_shape = ShapeFact::from_dims(dims);
|
||||
//
|
||||
// let res = ModelLoader
|
||||
// ::resolve_shape(&static_shape);
|
||||
// }
|
||||
}
|
||||
@@ -1 +1 @@
|
||||
pub mod session;
|
||||
pub mod session;
|
||||
|
||||
@@ -1,35 +1,61 @@
|
||||
use crate::loader::ModelLoader;
|
||||
use anyhow::Context;
|
||||
use crate::loader::ModelMetadata;
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::error::{DdddError, Result, TensorError};
|
||||
use ddddocr_core::utils::normalize_ocr_logits;
|
||||
use ddddocr_core::{InferenceEngine, ModelMetadata, OcrEngine, OcrOutput};
|
||||
use ndarray::s;
|
||||
use std::path::Path;
|
||||
use tract_onnx::prelude::DatumType;
|
||||
use tract_onnx::prelude::{Graph, IntoTensor, RunnableModel, Tensor, TypedFact, TypedOp, tvec};
|
||||
use ddddocr_core::{InferenceEngine, OcrEngine, OcrOutput};
|
||||
use tract_onnx::prelude::{DatumType, OutletId, ShapeFact, TypedModel};
|
||||
use tract_onnx::prelude::{IntoTensor, Tensor, tvec};
|
||||
// 引入核心层的统一错误类型
|
||||
/// 明确命名为 AxisDim,代表模型某一个轴的维度特征
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub enum AxisDim {
|
||||
/// 静态固定维度(如通道数固定为 1,高度固定为 64)
|
||||
Static(usize),
|
||||
/// 动态符号维度(如宽度是动态的 "image_width")
|
||||
Dynamic(String),
|
||||
}
|
||||
|
||||
impl AxisDim {
|
||||
/// 便捷方法:判断是否为动态维度
|
||||
pub fn is_dynamic(&self) -> bool {
|
||||
matches!(self, AxisDim::Dynamic(_))
|
||||
}
|
||||
}
|
||||
/// 自定义 Debug 格式化输出,彻底融化套娃外壳,保证日志干净漂亮
|
||||
impl std::fmt::Debug for AxisDim {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
AxisDim::Static(size) => write!(f, "{}", size),
|
||||
AxisDim::Dynamic(expr) => write!(f, "Dynamic(\"{}\")", expr),
|
||||
}
|
||||
}
|
||||
}
|
||||
/// 模拟 Python 的 input_info 和 output_info 结构
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TensorInfo {
|
||||
pub name: String,
|
||||
pub shape: Vec<AxisDim>, // 既包含 Fixed 静态维度,也包含 Dynamic 动态符号
|
||||
pub data_type: DatumType, // 对应 Python 的 type
|
||||
}
|
||||
|
||||
/// 最终返回的模型完整信息
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ModelInfo {
|
||||
pub inputs: Vec<TensorInfo>,
|
||||
pub outputs: Vec<TensorInfo>,
|
||||
/// 硬件执行提供者(采用 Option 兼容不同底层的推理引擎)
|
||||
pub providers: Option<Vec<String>>,
|
||||
}
|
||||
pub struct OcrSession {
|
||||
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
|
||||
pub session: Session,
|
||||
pub model_metadata: ModelMetadata,
|
||||
}
|
||||
impl OcrSession {
|
||||
pub fn new<P>(model_path: P, model_metadata: ModelMetadata) -> Result<Self>
|
||||
where
|
||||
P: AsRef<Path>,
|
||||
{
|
||||
let session = ModelLoader::model_for_path(model_path)?.session;
|
||||
Ok(Self {
|
||||
pub fn new(session: Session, model_metadata: ModelMetadata) -> Self {
|
||||
Self {
|
||||
session,
|
||||
model_metadata,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn model_from_bytes(model_bytes: &[u8], model_metadata: ModelMetadata) -> Result<Self> {
|
||||
let session = ModelLoader::model_from_bytes(model_bytes)?.session;
|
||||
Ok(Self {
|
||||
session,
|
||||
model_metadata,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
impl OcrEngine for OcrSession {
|
||||
@@ -40,7 +66,7 @@ impl OcrEngine for OcrSession {
|
||||
impl InferenceEngine for OcrSession {
|
||||
type Output = OcrOutput;
|
||||
/// 对应 Python 的 _inference
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output> {
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
|
||||
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
|
||||
// let result = self.ocr.run(tvec!(tensor.into()))?;
|
||||
let tensor = Tensor::from(input_array);
|
||||
@@ -48,12 +74,8 @@ impl InferenceEngine for OcrSession {
|
||||
let mut result = self
|
||||
.session
|
||||
.run(tvec!(tensor.into()))
|
||||
.map_err(|_| {
|
||||
DdddError::Inference(TensorError::EngineError(
|
||||
"执行模型推理失败".to_string(),
|
||||
))
|
||||
})?;
|
||||
// .context("执行模型推理失败")?;
|
||||
.map_err(|_| TensorError::Engine("执行模型推理失败".to_string()))?;
|
||||
// .context("执行模型推理失败")?;
|
||||
println!("模型输出原始数据: {:?}", result);
|
||||
// Ok(result.swap_remove(0).into_tensor())
|
||||
let raw_tensor = result.swap_remove(0).into_tensor();
|
||||
@@ -62,23 +84,17 @@ impl InferenceEngine for OcrSession {
|
||||
DatumType::I64 => {
|
||||
let array_d = raw_tensor
|
||||
.into_array::<i64>()
|
||||
.map_err(|_| {
|
||||
DdddError::Inference(TensorErrorReason::EngineError(
|
||||
"Tract 无法获取 i64 内存视图".to_string(),
|
||||
))
|
||||
})?;
|
||||
// .context("Tract 无法获取 i64 内存视图")?;
|
||||
.map_err(|_| TensorError::Engine("Tract 无法获取 i64 内存视图".to_string()))?;
|
||||
// .context("Tract 无法获取 i64 内存视图")?;
|
||||
// 🌟 提前提取真实维度
|
||||
let actual_shape = array_d.shape().to_vec();
|
||||
// 转成标准的 Array1 传给 core
|
||||
let array1 = array_d
|
||||
.to_owned()
|
||||
.into_dimensionality::<ndarray::Ix1>()
|
||||
.map_err(|_| {
|
||||
DdddError::Inference(TensorErrorReason::TensorDimensionMismatch {
|
||||
expected: "1D 字符索引静态矩阵".to_string(),
|
||||
actual: actual_shape,
|
||||
})
|
||||
.map_err(|_| TensorError::DimensionMismatch {
|
||||
expected: "1D 字符索引静态矩阵".to_string(),
|
||||
actual: actual_shape,
|
||||
})?;
|
||||
Ok(OcrOutput::Indices(array1))
|
||||
}
|
||||
@@ -87,18 +103,86 @@ impl InferenceEngine for OcrSession {
|
||||
println!("模型输出shape数据: {:?}", shape);
|
||||
let view = raw_tensor
|
||||
.to_array_view::<f32>()
|
||||
.map_err(|_| {
|
||||
DdddError::Inference(TensorErrorReason::EngineError(
|
||||
"Tract 无法获取 f32 内存视图".to_string(),
|
||||
))
|
||||
})?;
|
||||
.map_err(|_| TensorError::Engine("Tract 无法获取 f32 内存视图".to_string()))?;
|
||||
// 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗
|
||||
normalize_ocr_logits(view, shape)
|
||||
}
|
||||
_ => Err(
|
||||
// anyhow::anyhow!("不支持的模型输出数据类型: {:?}",raw_tensor.datum_type())
|
||||
DdddError::Inference(TensorErrorReason::UnknownOutputFormat)
|
||||
TensorError::UnknownOutputFormat,
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
impl OcrSession {
|
||||
/// 获取模型输入的节点信息列表
|
||||
pub fn input_info(&self) -> Result<Vec<TensorInfo>> {
|
||||
let model = self.session.model();
|
||||
let outlets = model.input_outlets().map_err(DdddError::new)?;
|
||||
self.resolve_tensors(model, outlets)
|
||||
}
|
||||
|
||||
/// 获取模型输出的节点信息列表
|
||||
pub fn output_info(&self) -> Result<Vec<TensorInfo>> {
|
||||
let model = self.session.model();
|
||||
let outlets = model.output_outlets().map_err(DdddError::new)?;
|
||||
self.resolve_tensors(model, outlets)
|
||||
}
|
||||
|
||||
/// 获取模型详细元数据信息(对标 Python ddddocr 的 get_model_info)
|
||||
/// 完美包容 [1, 1, 64, image_width] 这样的变长图像模型
|
||||
/// 获取模型详细元数据信息(代码更紧凑、优雅)
|
||||
pub fn model_info(&self) -> Result<ModelInfo> {
|
||||
Ok(ModelInfo {
|
||||
inputs: self.input_info()?,
|
||||
outputs: self.output_info()?,
|
||||
providers: None,
|
||||
})
|
||||
}
|
||||
|
||||
/// 提取出来的公共转换逻辑:将一组 OutletId 解析为 TensorInfo 列表
|
||||
fn resolve_tensors(&self, model: &TypedModel, outlets: &[OutletId]) -> Result<Vec<TensorInfo>> {
|
||||
outlets
|
||||
.iter()
|
||||
.map(|&outlet_id| {
|
||||
let fact = model.outlet_fact(outlet_id).map_err(DdddError::new)?;
|
||||
// .map_err(|e| {
|
||||
// DdddError::InternalError(format!("解析节点 Fact 失败: {:?}", e))
|
||||
// })?;
|
||||
|
||||
let shape = self.resolve_shape(&fact.shape);
|
||||
let node_name = model.node(outlet_id.node).name.clone();
|
||||
|
||||
Ok(TensorInfo {
|
||||
name: node_name,
|
||||
shape,
|
||||
data_type: fact.datum_type,
|
||||
})
|
||||
})
|
||||
.collect() // 函数式声明:自动传播第一处发生的错误
|
||||
}
|
||||
|
||||
/// 安全还原 Tract 维度至 Vec<AxisDim>
|
||||
fn resolve_shape(&self, shape_fact: &ShapeFact) -> Vec<AxisDim> {
|
||||
let tract_shape = shape_fact.to_tvec();
|
||||
|
||||
let resolved = tract_shape
|
||||
.iter()
|
||||
.map(|dim| {
|
||||
// 防御性编程:必须同时满足能够转换为 i64 且 大于等于 0
|
||||
if let Ok(size) = dim.to_i64() {
|
||||
if size >= 0 {
|
||||
AxisDim::Static(size as usize)
|
||||
} else {
|
||||
// 如果 ONNX 导出时某些动态维度被标记为了 -1,安全地作为动态符号捕获
|
||||
AxisDim::Dynamic(dim.to_string())
|
||||
}
|
||||
} else {
|
||||
AxisDim::Dynamic(dim.to_string())
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
resolved
|
||||
}
|
||||
}
|
||||
|
||||
3
ddddocr-tract/src/types.rs
Normal file
3
ddddocr-tract/src/types.rs
Normal file
@@ -0,0 +1,3 @@
|
||||
use tract_onnx::prelude::{Graph, RunnableModel, TypedFact, TypedOp};
|
||||
|
||||
pub type Session = RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>;
|
||||
@@ -2,8 +2,8 @@ use std::borrow::Cow;
|
||||
use std::fs::File;
|
||||
use std::path::Path;
|
||||
use anyhow::anyhow;
|
||||
use ddddocr_core::ocr::metadata::Charset;
|
||||
use ddddocr_core::ocr::metadata::{Normalization, Resize};
|
||||
use ddddocr_core::ocr::Charset;
|
||||
use ddddocr_core::ocr::{Normalization, Resize};
|
||||
|
||||
pub const CHARSET_BETA: &[&str] = &[
|
||||
"", "笤", "谴", "膀", "荔", "佰", "电", "臁", "矍", "同", "奇", "芄", "吠", "6", "曛", "荇",
|
||||
|
||||
@@ -1,15 +1,17 @@
|
||||
use anyhow::Context;
|
||||
use ddddocr_core::det::DetectionResult;
|
||||
use ddddocr_core::{DetBuilder, Detector, ModelMetadata, Ocr, Slider}; // 假设你的包名是这个
|
||||
use ddddocr_tract::{DetSession, OcrSession};
|
||||
use ddddocr_core::{Detector, ModelMetadata, Normalization, Ocr, Slider};
|
||||
// 假设你的包名是这个
|
||||
use ddddocr_tract::{DetSession, OcrSession,OcrBuilder};
|
||||
use image::{DynamicImage, ImageBuffer, Luma, Rgb};
|
||||
use std::fs;
|
||||
use std::path::Path;
|
||||
use tract_onnx::prelude::{ShapeFact, TDim};
|
||||
use tract_onnx::model;
|
||||
|
||||
mod char_slice;
|
||||
use char_slice::CHARSET_BETA;
|
||||
use ddddocr_core::ocr::metadata::{Normalization, Resize};
|
||||
use ddddocr_core::ocr::Resize;
|
||||
|
||||
use ddddocr_tract::loader::ModelLoader;
|
||||
|
||||
fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
|
||||
@@ -103,25 +105,29 @@ fn save_rust_result(result: &ImageBuffer<Luma<f32>, Vec<f32>>, filename: &str) {
|
||||
}
|
||||
#[test]
|
||||
fn test_full_classification() {
|
||||
// 1. 初始化模型
|
||||
let ocr = OcrSession::new(
|
||||
let model = ModelLoader::model_for_path(
|
||||
"D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx",
|
||||
ModelMetadata::from_static_slice(
|
||||
CHARSET_BETA,
|
||||
false,
|
||||
Resize::DynamicWidth(64),
|
||||
1,
|
||||
Normalization::MinusOneToOne,
|
||||
),
|
||||
)
|
||||
.expect("模型加载失败");
|
||||
|
||||
let metadata = ModelMetadata::from_static_slice(
|
||||
CHARSET_BETA,
|
||||
false,
|
||||
Resize::DynamicWidth(64),
|
||||
1,
|
||||
Normalization::MinusOneToOne,
|
||||
);
|
||||
// 1. 初始化模型
|
||||
let ocr = OcrSession::new(model, metadata);
|
||||
// 2. 加载测试图片
|
||||
let img =
|
||||
image::open("D:/CNWei/CNW/Rust/ddddocr-rs/samples/code2.png").expect("测试图片不存在");
|
||||
|
||||
// 3. 执行识别
|
||||
let result = Ocr::new(&ocr)
|
||||
// let result = Ocr::new(&ocr)
|
||||
// .predict(&img)
|
||||
// .expect("识别过程出错")
|
||||
// .into_text();
|
||||
let result = OcrBuilder::new().build(&ocr)
|
||||
.predict(&img)
|
||||
.expect("识别过程出错")
|
||||
.into_text();
|
||||
@@ -131,7 +137,10 @@ fn test_full_classification() {
|
||||
}
|
||||
#[test]
|
||||
fn test_det_load() -> anyhow::Result<()> {
|
||||
let det = DetSession::new("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_det.onnx")?;
|
||||
let det_model =
|
||||
ModelLoader::model_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_det.onnx")
|
||||
.expect("模型加载失败");
|
||||
let det = DetSession::new(det_model);
|
||||
let image_path = "D:/CNWei/CNW/Rust/ddddocr-rs/samples/det1.png";
|
||||
let image_bytes =
|
||||
fs::read(image_path).map_err(|e| anyhow::anyhow!("无法读取图片 {}: {}", image_path, e))?;
|
||||
@@ -167,7 +176,7 @@ fn test_det_load() -> anyhow::Result<()> {
|
||||
|
||||
#[test]
|
||||
fn test_real_slide_match() {
|
||||
let engine = Slider::new().unwrap();
|
||||
let engine = Slider::new();
|
||||
|
||||
// 1. 加载你准备好的测试图
|
||||
// 假设图片放在项目根目录下的 assets 文件夹
|
||||
@@ -198,7 +207,7 @@ fn test_real_slide_match() {
|
||||
|
||||
#[test]
|
||||
fn test_real_slide_comparison() {
|
||||
let engine = Slider::new().unwrap();
|
||||
let engine = Slider::new();
|
||||
|
||||
// 1. 加载你准备好的测试图
|
||||
// 假设图片放在项目根目录下的 assets 文件夹
|
||||
@@ -236,6 +245,4 @@ fn test_resolve_shape_logic_direct() {
|
||||
"D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_huashi666_i64.onnx",
|
||||
)
|
||||
.expect("建立测试模型图失败");
|
||||
let md_info = &loader.model_info().context("信息");
|
||||
println!("md_info: {:?}", md_info);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user