refactor: 优化 ModelLoader 结构与维度解析逻辑

- 提取 `resolve_tensors` 公共函数,消除输入输出流的重复代码。
- 简化 `resolve_shape`,使其回归纯粹的维度翻译职责,移除过早的错误校验。
- 统一错误处理,将底层解析异常清晰映射至 `DdddError::InternalError`。
This commit is contained in:
2026-07-14 18:56:41 +08:00
parent 4fd38022fd
commit cd70748393
3 changed files with 188 additions and 4 deletions

View File

@@ -17,7 +17,6 @@ pub(crate) const MODEL_DOWNLOAD_HELP: &str = "\
B. 或者直接将模型文件重命名并放置在您运行程序的“当前工作目录”或“可执行文件同级目录”下。
================================================================================";
use thiserror::Error;
#[derive(Error, Debug)]
@@ -42,7 +41,10 @@ pub enum DdddError {
#[error("未知的模型输出格式")]
UnknownOutputFormat,
#[error("解析节点 Fact 失败")]
InternalError(String),
}
/// 统一用我们自己的 DdddError 包装 Result
pub type Result<T> = std::result::Result<T, DdddError>;
pub type Result<T> = std::result::Result<T, DdddError>;

View File

@@ -1,8 +1,50 @@
use anyhow::Context;
use ddddocr_core::error::Result;
use ddddocr_core::error::{DdddError, Result};
use std::fmt;
use std::io::Cursor;
use tract_onnx::onnx;
use tract_onnx::prelude::*; // 引入核心层的统一错误类型
use tract_onnx::prelude::*;
// 引入核心层的统一错误类型
/// 明确命名为 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 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>>>,
@@ -38,3 +80,132 @@ impl ModelLoader {
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(|e| DdddError::InternalError(format!("获取输入节点失败: {:?}", e)))?,
)?;
let outputs = self.resolve_tensors(
model
.output_outlets()
.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(|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 失败");
}
}

View File

@@ -4,9 +4,13 @@ use ddddocr_tract::{DetSession,OcrSession};
use image::{DynamicImage, Rgb};
use std::fs;
use std::path::Path;
use anyhow::Context;
use tract_onnx::prelude::{ShapeFact, TDim};
mod char_slice;
use char_slice::CHARSET_BETA;
use ddddocr_core::ocr::metadata::{Normalization, Resize};
use ddddocr_tract::loader::ModelLoader;
fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
// 1. 先将泛型转为具体的 &Path 引用
@@ -182,3 +186,10 @@ fn test_real_slide_comparison() {
assert_eq!(result.target_y, 90);
assert!(result.confidence > 0.0);
}
#[test]
fn test_resolve_shape_logic_direct() {
// 创建一个哑 ModelLoader 实例session 用不上,因为我们直接测私有方法)
let loader = ModelLoader::model_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx",).expect("建立测试模型图失败");
let md_info=&loader.model_info().context("信息");
println!("md_info: {:?}",md_info);
}