diff --git a/ddddocr-core/src/error.rs b/ddddocr-core/src/error.rs index f81fbea..529ad7f 100644 --- a/ddddocr-core/src/error.rs +++ b/ddddocr-core/src/error.rs @@ -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 = std::result::Result; \ No newline at end of file +pub type Result = std::result::Result; diff --git a/ddddocr-tract/src/loader.rs b/ddddocr-tract/src/loader.rs index 762b039..e4431bd 100644 --- a/ddddocr-tract/src/loader.rs +++ b/ddddocr-tract/src/loader.rs @@ -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, // 既包含 Fixed 静态维度,也包含 Dynamic 动态符号 + pub data_type: DatumType, // 对应 Python 的 type +} + +/// 最终返回的模型完整信息 +#[derive(Debug, Clone)] +pub struct ModelInfo { + pub inputs: Vec, + pub outputs: Vec, + /// 硬件执行提供者(采用 Option 兼容不同底层的推理引擎) + pub providers: Option>, +} pub struct ModelLoader { pub session: RunnableModel, Graph>>, @@ -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 { + 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> { + 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 + fn resolve_shape(&self, shape_fact: &ShapeFact) -> Result> { + 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, Graph>>, + 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 = 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 失败"); + } +} diff --git a/ddddocr-tract/tests/ocr_test.rs b/ddddocr-tract/tests/ocr_test.rs index 2409230..fe0c9bb 100644 --- a/ddddocr-tract/tests/ocr_test.rs +++ b/ddddocr-tract/tests/ocr_test.rs @@ -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>(path: P) -> anyhow::Result { // 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); +} \ No newline at end of file