feat: ddddocr-rs 完成 core/ort/tract2 规范整改与发布
准备 - core:内置官方字符集(OLD/BETA)与 ModelMetadata::from_builtin_* 构造器 - ort:导出 Session、实现 Info、修正 cuda feature 接线、共享推理工具 - tract2:包名更名(原 ddddocr-tract 已被占用)并完成规范整改 - 集成测试按领域拆分(ocr / det / slide / common / api_surface) - 发布准备:Cargo.toml 元数据、workspace 版本 0.2.4、LICENSE/ NOTICE、README 模型下载说明
This commit is contained in:
22
ddddocr-tract2/Cargo.toml
Normal file
22
ddddocr-tract2/Cargo.toml
Normal file
@@ -0,0 +1,22 @@
|
||||
[package]
|
||||
name = "ddddocr-tract2"
|
||||
version = { workspace = true }
|
||||
edition = { workspace = true }
|
||||
license = { workspace = true }
|
||||
description = "ddddocr-rs 的 Tract 推理引擎实现"
|
||||
keywords = ["ocr", "captcha", "ddddocr", "onnxruntime", "tract"]
|
||||
categories = ["multimedia::images", "computer-vision"]
|
||||
# repository = "https://github.com/<用户名>/<仓库名>" # 发布前请补充
|
||||
readme = "../README.md"
|
||||
|
||||
[dependencies]
|
||||
ddddocr-core = { path = "../ddddocr-core", version = "0.2.4" }
|
||||
tract-onnx = { workspace = true }
|
||||
tract-linalg = { workspace = true }
|
||||
ndarray = { workspace = true }
|
||||
image = { workspace = true }
|
||||
serde = { workspace = true }
|
||||
serde_json = { workspace = true }
|
||||
anyhow = { workspace = true }
|
||||
thiserror = { workspace = true }
|
||||
tracing = { workspace = true }
|
||||
1
ddddocr-tract2/src/det.rs
Normal file
1
ddddocr-tract2/src/det.rs
Normal file
@@ -0,0 +1 @@
|
||||
pub mod session;
|
||||
44
ddddocr-tract2/src/det/session.rs
Normal file
44
ddddocr-tract2/src/det/session.rs
Normal file
@@ -0,0 +1,44 @@
|
||||
use crate::runtime::{extract_plain_array, run_session};
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::DetOutput;
|
||||
use ddddocr_core::error::{Result, TensorError};
|
||||
use ddddocr_core::traits::{DetEngine, InferenceEngine};
|
||||
use ndarray::Ix3;
|
||||
use tracing::debug;
|
||||
use tract_onnx::prelude::*;
|
||||
|
||||
#[derive(Debug)]
|
||||
/// 目标检测推理运行时:持有 Tract 会话,输出 [`DetOutput`]。
|
||||
pub struct DetRuntime {
|
||||
/// Tract 会话句柄。
|
||||
pub session: Session,
|
||||
}
|
||||
|
||||
impl DetRuntime {
|
||||
/// 基于已构建的会话创建检测运行时。
|
||||
pub fn new(session: Session) -> Self {
|
||||
Self { session }
|
||||
}
|
||||
}
|
||||
|
||||
impl InferenceEngine for DetRuntime {
|
||||
type Output = DetOutput;
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
|
||||
let mut result = run_session(&self.session, input_array)?;
|
||||
debug!("模型输出原始数据: {:?}", result);
|
||||
let raw_tensor = result.swap_remove(0).into_tensor();
|
||||
let array_d = extract_plain_array::<f32>(raw_tensor)?;
|
||||
let actual_shape = array_d.shape().to_vec();
|
||||
|
||||
let array3 =
|
||||
array_d
|
||||
.into_dimensionality::<Ix3>()
|
||||
.map_err(|_| TensorError::DimensionMismatch {
|
||||
expected: "3D 检测矩阵 [Batch, Box_Count, Box_Attributes]".to_string(),
|
||||
actual: actual_shape,
|
||||
})?;
|
||||
Ok(DetOutput::Detection(array3))
|
||||
}
|
||||
}
|
||||
|
||||
impl DetEngine for DetRuntime {}
|
||||
22
ddddocr-tract2/src/lib.rs
Normal file
22
ddddocr-tract2/src/lib.rs
Normal file
@@ -0,0 +1,22 @@
|
||||
//! # ddddocr-tract2
|
||||
//!
|
||||
//! 基于 Tract(纯 Rust ONNX 推理框架)的推理引擎实现:通过 [`loader::ModelLoader`] 构建会话,
|
||||
//! 再由 [`OcrRuntime`] / [`DetRuntime`] 实现 core 的
|
||||
//! [`ddddocr_core::traits::InferenceEngine`] 接口。
|
||||
//! 同时重导出 core 的 [`OcrBuilder`]、[`Slider`]、[`SlideResult`] 等便捷 API。
|
||||
|
||||
#![warn(missing_docs)]
|
||||
|
||||
mod det;
|
||||
/// 模型加载器:从路径或字节流构建 Tract 会话。
|
||||
pub mod loader;
|
||||
mod ocr;
|
||||
mod runtime;
|
||||
mod types;
|
||||
|
||||
pub use ddddocr_core::{
|
||||
DetectionResult, Detector, ModelMetadata, Normalization, Ocr, OcrBuilder, SlideResult, Slider,
|
||||
};
|
||||
pub use det::session::DetRuntime;
|
||||
pub use ocr::session::OcrRuntime;
|
||||
pub use types::Session;
|
||||
7
ddddocr-tract2/src/loader.rs
Normal file
7
ddddocr-tract2/src/loader.rs
Normal file
@@ -0,0 +1,7 @@
|
||||
mod error;
|
||||
mod metadata;
|
||||
mod model;
|
||||
|
||||
pub use error::{Error, ParseError, Result};
|
||||
pub use metadata::{Metadata, ModelMetadataDto, NormalizationDto};
|
||||
pub use model::ModelLoader;
|
||||
82
ddddocr-tract2/src/loader/error.rs
Normal file
82
ddddocr-tract2/src/loader/error.rs
Normal file
@@ -0,0 +1,82 @@
|
||||
use tract_onnx::prelude::TractError;
|
||||
/// 模型加载与解析的通用结果类型。
|
||||
pub type Result<T> = std::result::Result<T, Error>;
|
||||
#[derive(thiserror::Error, Debug)]
|
||||
/// 模型加载、解析与 Session 构建阶段的错误。
|
||||
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)]
|
||||
/// 从路径或字节流解析 ONNX 模型失败的错误。
|
||||
pub enum ParseError {
|
||||
/// 策略 A:从文件路径加载失败(附带路径上下文信息,方便排查是找不到文件还是格式不对)
|
||||
#[error("从路径 '{0}' 加载 ONNX 模型失败: {1}")]
|
||||
Path(String, #[source] TractError),
|
||||
|
||||
/// 策略 B:从内存字节流加载失败(如 include_bytes! 传入的字节流损坏)
|
||||
#[error("从内存字节流解析 ONNX 模型失败: {0}")]
|
||||
Bytes(#[source] TractError),
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn wraps_external_error_via_new() {
|
||||
let io_err = std::io::Error::other("boom");
|
||||
let err = Error::new("自定义错误", io_err);
|
||||
assert!(matches!(err, Error::Other(_, _)));
|
||||
assert!(err.to_string().contains("自定义错误"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn serde_json_error_converts_to_error() {
|
||||
let json_err = serde_json::from_str::<serde_json::Value>("{").unwrap_err();
|
||||
let err: Error = json_err.into();
|
||||
assert!(matches!(err, Error::JsonParse(_)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_utf8_converts_to_error() {
|
||||
let bytes = vec![0xffu8];
|
||||
let utf8_err = std::str::from_utf8(&bytes).unwrap_err();
|
||||
let err: Error = utf8_err.into();
|
||||
assert!(matches!(err, Error::InvalidUtf8(_)));
|
||||
}
|
||||
}
|
||||
195
ddddocr-tract2/src/loader/metadata.rs
Normal file
195
ddddocr-tract2/src/loader/metadata.rs
Normal file
@@ -0,0 +1,195 @@
|
||||
use crate::loader::error::{Error, Result};
|
||||
|
||||
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"
|
||||
/// 归一化策略的 JSON 反序列化中间表示。
|
||||
pub 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)]
|
||||
pub 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
|
||||
}
|
||||
|
||||
/// 从 JSON 字符串或字节流解析模型元数据的扩展接口。
|
||||
pub trait Metadata: Sized {
|
||||
/// 从 JSON 字符串解析模型元数据。
|
||||
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 Metadata 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(Cow::Owned).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(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
const MINIMAL_JSON: &str = r#"{
|
||||
"charset": ["a", "b", "c"],
|
||||
"word": false,
|
||||
"resize": [-1, 64],
|
||||
"channel": 1
|
||||
}"#;
|
||||
|
||||
#[test]
|
||||
fn parses_minimal_json_with_default_normalization() {
|
||||
let meta = ModelMetadata::from_json_str(MINIMAL_JSON).unwrap();
|
||||
assert_eq!(meta.charset.size(), 3);
|
||||
assert!(!meta.word);
|
||||
assert_eq!(meta.channel, 1);
|
||||
assert!(matches!(meta.resize, Resize::DynamicWidth(64)));
|
||||
assert!(matches!(meta.normalization, Normalization::ZeroToOne));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_minus_one_to_one_normalization() {
|
||||
let json = r#"{
|
||||
"charset": ["a"],
|
||||
"word": false,
|
||||
"resize": [-1, 64],
|
||||
"channel": 1,
|
||||
"normalization": "minus_one_to_one"
|
||||
}"#;
|
||||
let meta = ModelMetadata::from_json_str(json).unwrap();
|
||||
assert!(matches!(meta.normalization, Normalization::MinusOneToOne));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_word_model_as_square_resize() {
|
||||
let json = r#"{
|
||||
"charset": ["a"],
|
||||
"word": true,
|
||||
"resize": [-1, 64],
|
||||
"channel": 1
|
||||
}"#;
|
||||
let meta = ModelMetadata::from_json_str(json).unwrap();
|
||||
assert!(matches!(meta.resize, Resize::Square(64)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_fixed_resize() {
|
||||
let json = r#"{
|
||||
"charset": ["a"],
|
||||
"word": false,
|
||||
"resize": [100, 64],
|
||||
"channel": 1
|
||||
}"#;
|
||||
let meta = ModelMetadata::from_json_str(json).unwrap();
|
||||
assert!(matches!(meta.resize, Resize::Fixed(100, 64)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_image_alias() {
|
||||
let json = r#"{
|
||||
"charset": ["a"],
|
||||
"word": false,
|
||||
"image": [200, 64],
|
||||
"channel": 1
|
||||
}"#;
|
||||
let meta = ModelMetadata::from_json_str(json).unwrap();
|
||||
assert!(matches!(meta.resize, Resize::Fixed(200, 64)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_resize_with_wrong_length() {
|
||||
let json = r#"{
|
||||
"charset": ["a"],
|
||||
"word": false,
|
||||
"resize": [1, 2, 3],
|
||||
"channel": 1
|
||||
}"#;
|
||||
assert!(matches!(
|
||||
ModelMetadata::from_json_str(json),
|
||||
Err(Error::MetadataParse(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_json() {
|
||||
assert!(matches!(
|
||||
ModelMetadata::from_json_str("not json"),
|
||||
Err(Error::JsonParse(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_utf8_bytes() {
|
||||
assert!(matches!(
|
||||
ModelMetadata::from_json_bytes(&[0xff, 0xfe]),
|
||||
Err(Error::InvalidUtf8(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
95
ddddocr-tract2/src/loader/model.rs
Normal file
95
ddddocr-tract2/src/loader/model.rs
Normal file
@@ -0,0 +1,95 @@
|
||||
use crate::loader::error::{Error, ParseError, Result};
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::traits::Loader;
|
||||
use std::io::Cursor;
|
||||
use tract_linalg::multithread::{Executor, set_default_executor};
|
||||
use tract_onnx::onnx;
|
||||
use tract_onnx::prelude::*;
|
||||
|
||||
/// Tract 专用的链式构建器
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct ModelLoader {
|
||||
num_threads: Option<usize>,
|
||||
}
|
||||
|
||||
impl ModelLoader {
|
||||
/// 可选扩展:设置 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 Loader for ModelLoader {
|
||||
type Session = Session;
|
||||
type Error = Error;
|
||||
fn build_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()
|
||||
.model_for_path(path_ref)
|
||||
.map_err(|e| ParseError::Path(path_ref.display().to_string(), e))?
|
||||
.into_optimized()
|
||||
.map_err(Error::OptimizationFailed)?
|
||||
.into_runnable()
|
||||
.map_err(Error::RunnableBuildFailed)?;
|
||||
Ok(session)
|
||||
}
|
||||
/// 策略 B:从内存字节流加载模型(配合 include_bytes! 使用)
|
||||
fn build_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);
|
||||
|
||||
let session = onnx()
|
||||
.model_for_read(&mut cursor)
|
||||
.map_err(ParseError::Bytes)?
|
||||
.into_optimized()
|
||||
.map_err(Error::OptimizationFailed)?
|
||||
.into_runnable()
|
||||
.map_err(Error::RunnableBuildFailed)?;
|
||||
|
||||
Ok(session)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn default_config() {
|
||||
let loader = ModelLoader::default();
|
||||
assert!(loader.num_threads.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builder_configures_threads() {
|
||||
let loader = ModelLoader::default().num_threads(4);
|
||||
assert_eq!(loader.num_threads, Some(4));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builder_is_chainable_and_immutable() {
|
||||
let base = ModelLoader::default();
|
||||
let _configured = base.clone().num_threads(8);
|
||||
assert!(base.num_threads.is_none());
|
||||
}
|
||||
}
|
||||
1
ddddocr-tract2/src/ocr.rs
Normal file
1
ddddocr-tract2/src/ocr.rs
Normal file
@@ -0,0 +1 @@
|
||||
pub mod session;
|
||||
174
ddddocr-tract2/src/ocr/session.rs
Normal file
174
ddddocr-tract2/src/ocr/session.rs
Normal file
@@ -0,0 +1,174 @@
|
||||
use crate::runtime::{extract_plain_array, run_session};
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::ModelMetadata;
|
||||
use ddddocr_core::OcrOutput;
|
||||
use ddddocr_core::error::{DdddError, Result, TensorError};
|
||||
use ddddocr_core::traits::{InferenceEngine, Info, OcrEngine};
|
||||
use ddddocr_core::types::{AxisDim, ModelInfo, TensorInfo, TensorType};
|
||||
use ddddocr_core::utils::normalize_ocr_logits;
|
||||
use tracing::debug;
|
||||
use tract_onnx::prelude::{DatumType, IntoTensor, OutletId, ShapeFact, TypedModel};
|
||||
|
||||
/// OCR 推理运行时:持有 Tract 会话与模型元数据,输出 [`OcrOutput`]。
|
||||
pub struct OcrRuntime {
|
||||
/// Tract 会话句柄。
|
||||
pub session: Session,
|
||||
/// 模型元数据(字符集、归一化策略等)。
|
||||
pub metadata: ModelMetadata,
|
||||
}
|
||||
impl OcrRuntime {
|
||||
/// 基于已构建的会话与元数据创建 OCR 运行时。
|
||||
pub fn new(session: Session, metadata: ModelMetadata) -> Self {
|
||||
Self { session, metadata }
|
||||
}
|
||||
|
||||
/// 提取出来的公共转换逻辑:将一组 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)?;
|
||||
let shape = resolve_shape(&fact.shape);
|
||||
let node_name = model.node(outlet_id.node).name.clone();
|
||||
let tensor_type = tensor_type_from_datum(fact.datum_type);
|
||||
|
||||
Ok(TensorInfo {
|
||||
name: node_name,
|
||||
shape,
|
||||
tensor_type,
|
||||
})
|
||||
})
|
||||
.collect() // 函数式声明:自动传播第一处发生的错误
|
||||
}
|
||||
}
|
||||
|
||||
/// 将 Tract 的 DatumType 映射为 core 的 [`TensorType`]。
|
||||
fn tensor_type_from_datum(datum: DatumType) -> TensorType {
|
||||
match datum {
|
||||
DatumType::F32 => TensorType::F32,
|
||||
DatumType::I64 => TensorType::I64,
|
||||
_ => TensorType::Other,
|
||||
}
|
||||
}
|
||||
|
||||
/// 安全还原 Tract 维度至 [`AxisDim`] 列表。
|
||||
fn resolve_shape(shape_fact: &ShapeFact) -> Vec<AxisDim> {
|
||||
shape_fact
|
||||
.to_tvec()
|
||||
.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()
|
||||
}
|
||||
|
||||
impl OcrEngine for OcrRuntime {
|
||||
fn metadata(&self) -> &ModelMetadata {
|
||||
&self.metadata
|
||||
}
|
||||
}
|
||||
impl InferenceEngine for OcrRuntime {
|
||||
type Output = OcrOutput;
|
||||
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
|
||||
let mut result = run_session(&self.session, input_array)?;
|
||||
debug!("模型输出原始数据: {:?}", result);
|
||||
let raw_tensor = result.swap_remove(0).into_tensor();
|
||||
match raw_tensor.datum_type() {
|
||||
DatumType::I64 => {
|
||||
let array_d = extract_plain_array::<i64>(raw_tensor)?;
|
||||
let actual_shape = array_d.shape().to_vec();
|
||||
let array1 = array_d.into_dimensionality::<ndarray::Ix1>().map_err(|_| {
|
||||
TensorError::DimensionMismatch {
|
||||
expected: "1D 字符索引静态矩阵".to_string(),
|
||||
actual: actual_shape,
|
||||
}
|
||||
})?;
|
||||
Ok(OcrOutput::Indices(array1))
|
||||
}
|
||||
DatumType::F32 => {
|
||||
let shape = raw_tensor.shape().to_vec();
|
||||
debug!("模型输出 shape 数据: {:?}", shape);
|
||||
let array = extract_plain_array::<f32>(raw_tensor)?;
|
||||
normalize_ocr_logits(array.view(), array.shape())
|
||||
}
|
||||
_ => Err(TensorError::UnknownOutputFormat),
|
||||
}
|
||||
}
|
||||
}
|
||||
impl Info for OcrRuntime {
|
||||
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)
|
||||
}
|
||||
|
||||
/// 获取模型输出的节点信息列表
|
||||
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] 这样的变长图像模型
|
||||
/// 获取模型详细元数据信息(代码更紧凑、优雅)
|
||||
fn model_info(&self) -> Result<ModelInfo> {
|
||||
Ok(ModelInfo {
|
||||
inputs: self.input_info()?,
|
||||
outputs: self.output_info()?,
|
||||
providers: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use tract_onnx::prelude::*;
|
||||
|
||||
#[test]
|
||||
fn maps_datum_types() {
|
||||
assert!(matches!(
|
||||
tensor_type_from_datum(DatumType::F32),
|
||||
TensorType::F32
|
||||
));
|
||||
assert!(matches!(
|
||||
tensor_type_from_datum(DatumType::I64),
|
||||
TensorType::I64
|
||||
));
|
||||
assert!(matches!(
|
||||
tensor_type_from_datum(DatumType::U8),
|
||||
TensorType::Other
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_static_shape() {
|
||||
let fact = ShapeFact::from_dims(vec![TDim::from(1), TDim::from(64)]);
|
||||
assert_eq!(
|
||||
resolve_shape(&fact),
|
||||
vec![AxisDim::Static(1), AxisDim::Static(64)]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_dynamic_dimension_as_symbol() {
|
||||
let scope = SymbolScope::default();
|
||||
let dims = vec![TDim::from(1), TDim::from(64), scope.sym("width").into()];
|
||||
let fact = ShapeFact::from_dims(dims);
|
||||
let shape = resolve_shape(&fact);
|
||||
assert_eq!(shape[0], AxisDim::Static(1));
|
||||
assert_eq!(shape[1], AxisDim::Static(64));
|
||||
assert!(matches!(&shape[2], AxisDim::Dynamic(_)));
|
||||
}
|
||||
}
|
||||
25
ddddocr-tract2/src/runtime.rs
Normal file
25
ddddocr-tract2/src/runtime.rs
Normal file
@@ -0,0 +1,25 @@
|
||||
//! 会话执行相关的共享工具函数。
|
||||
|
||||
use crate::types::Session;
|
||||
use ddddocr_core::error::TensorError;
|
||||
use tract_onnx::prelude::{Datum, TValue, TVec, Tensor, tvec};
|
||||
|
||||
/// 对输入张量执行一次 Tract 推理,返回原始输出值列表。
|
||||
pub(crate) fn run_session(
|
||||
session: &Session,
|
||||
input_array: ndarray::Array4<f32>,
|
||||
) -> Result<TVec<TValue>, TensorError> {
|
||||
let tensor = Tensor::from(input_array);
|
||||
session
|
||||
.run(tvec!(tensor.into()))
|
||||
.map_err(|e| TensorError::Engine(format!("执行模型推理失败: {e}")))
|
||||
}
|
||||
|
||||
/// 将 Tensor 转换为 ndarray 的 `ArrayD`。
|
||||
pub(crate) fn extract_plain_array<D: Datum>(
|
||||
tensor: Tensor,
|
||||
) -> Result<ndarray::ArrayD<D>, TensorError> {
|
||||
tensor
|
||||
.into_plain_array::<D>()
|
||||
.map_err(|_| TensorError::Engine("无法获取张量内存视图".to_string()))
|
||||
}
|
||||
7
ddddocr-tract2/src/types.rs
Normal file
7
ddddocr-tract2/src/types.rs
Normal file
@@ -0,0 +1,7 @@
|
||||
//! Tract 会话共享类型。
|
||||
|
||||
use std::sync::Arc;
|
||||
use tract_onnx::prelude::TypedRunnableModel;
|
||||
|
||||
/// Tract 会话句柄:由 [`crate::loader::ModelLoader`] 构建,可直接并发执行推理。
|
||||
pub type Session = Arc<TypedRunnableModel>;
|
||||
74
ddddocr-tract2/tests/api_surface.rs
Normal file
74
ddddocr-tract2/tests/api_surface.rs
Normal file
@@ -0,0 +1,74 @@
|
||||
//! 外部视角 API 测试:验证 `ddddocr-tract2` 的公开类型与方法在外部 crate 中可正常使用。
|
||||
|
||||
use ddddocr_core::Resize;
|
||||
use ddddocr_core::traits::{Info, Loader};
|
||||
use ddddocr_core::types::AxisDim;
|
||||
use ddddocr_tract2::loader::ModelLoader;
|
||||
use ddddocr_tract2::{
|
||||
DetRuntime, ModelMetadata, Normalization, OcrBuilder, OcrRuntime, Session, SlideResult, Slider,
|
||||
};
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
fn model_path(name: &str) -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("..")
|
||||
.join("models")
|
||||
.join(name)
|
||||
}
|
||||
|
||||
/// 验证对外导出的类型(含 `Session` 与 core 便捷重导出)均可直接命名。
|
||||
#[test]
|
||||
fn public_types_are_nameable() {
|
||||
let _: Option<Session> = None;
|
||||
let _: Option<OcrBuilder> = None;
|
||||
let _: Option<SlideResult> = None;
|
||||
let _slider = Slider::new();
|
||||
}
|
||||
|
||||
/// 验证构建器链式 API 可组合使用,且返回的会话类型可显式标注。
|
||||
#[test]
|
||||
fn loader_chain_builds_session() -> anyhow::Result<()> {
|
||||
let path = model_path("common_sml2h3_f32.onnx");
|
||||
assert!(path.exists(), "缺少测试模型: {}", path.display());
|
||||
let _session: Session = ModelLoader::default().num_threads(4).build_for_path(path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 验证 `Info` trait 能从真实会话中解析输入/输出信息。
|
||||
#[test]
|
||||
fn info_trait_returns_model_metadata() -> anyhow::Result<()> {
|
||||
let path = model_path("common_sml2h3_f32.onnx");
|
||||
assert!(path.exists(), "缺少测试模型: {}", path.display());
|
||||
let session: Session = ModelLoader::default().build_for_path(path)?;
|
||||
let metadata = ModelMetadata::from_static_slice(
|
||||
&["a", "b"],
|
||||
false,
|
||||
Resize::DynamicWidth(64),
|
||||
1,
|
||||
Normalization::MinusOneToOne,
|
||||
);
|
||||
let ocr = OcrRuntime::new(session, metadata);
|
||||
|
||||
let inputs = ocr.input_info()?;
|
||||
let outputs = ocr.output_info()?;
|
||||
assert!(!inputs.is_empty());
|
||||
assert!(!outputs.is_empty());
|
||||
assert!(matches!(inputs[0].shape[0], AxisDim::Static(1)));
|
||||
assert!(matches!(inputs[0].shape[2], AxisDim::Static(64)));
|
||||
|
||||
let model_info = ocr.model_info()?;
|
||||
assert_eq!(model_info.inputs.len(), inputs.len());
|
||||
assert_eq!(model_info.outputs.len(), outputs.len());
|
||||
assert!(model_info.providers.is_none());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// 验证检测运行时可以从构建的会话创建。
|
||||
#[test]
|
||||
fn det_runtime_builds_from_session() -> anyhow::Result<()> {
|
||||
let path = model_path("common_det.onnx");
|
||||
assert!(path.exists(), "缺少测试模型: {}", path.display());
|
||||
let session: Session = ModelLoader::default().build_for_path(path)?;
|
||||
let _det = DetRuntime::new(session);
|
||||
Ok(())
|
||||
}
|
||||
25
ddddocr-tract2/tests/common/mod.rs
Normal file
25
ddddocr-tract2/tests/common/mod.rs
Normal file
@@ -0,0 +1,25 @@
|
||||
//! 集成测试共享工具。
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
/// 仓库根目录 `models/` 下模型文件的路径。
|
||||
pub fn model_path(name: &str) -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("..")
|
||||
.join("models")
|
||||
.join(name)
|
||||
}
|
||||
|
||||
/// 仓库根目录 `samples/` 下样例图片的路径。
|
||||
pub fn sample_path(name: &str) -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("..")
|
||||
.join("samples")
|
||||
.join(name)
|
||||
}
|
||||
|
||||
/// 加载图片,失败时附带路径上下文。
|
||||
pub fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
|
||||
let path_ref = path.as_ref();
|
||||
image::open(path_ref).map_err(|e| anyhow::anyhow!("无法加载图片 {:?}: {}", path_ref, e))
|
||||
}
|
||||
33
ddddocr-tract2/tests/det.rs
Normal file
33
ddddocr-tract2/tests/det.rs
Normal file
@@ -0,0 +1,33 @@
|
||||
//! 目标检测集成测试。
|
||||
|
||||
mod common;
|
||||
|
||||
use common::{model_path, sample_path};
|
||||
use ddddocr_core::Detector;
|
||||
use ddddocr_core::traits::Loader;
|
||||
use ddddocr_tract2::DetRuntime;
|
||||
use ddddocr_tract2::loader::ModelLoader;
|
||||
use image::GenericImageView;
|
||||
|
||||
/// 检测模型应能从样例图片中找到至少一个目标,且坐标在图片范围内。
|
||||
#[test]
|
||||
fn det_model_detects_targets_in_image() -> anyhow::Result<()> {
|
||||
let session = ModelLoader::default()
|
||||
.build_for_path(model_path("common_det.onnx"))
|
||||
.expect("模型加载失败");
|
||||
let det = DetRuntime::new(session);
|
||||
let img = image::open(sample_path("det1.png")).expect("测试图片不存在");
|
||||
|
||||
let bboxes = Detector::new(&det).predict(&img)?;
|
||||
assert!(!bboxes.is_empty(), "应检测到至少一个目标");
|
||||
|
||||
let (width, height) = img.dimensions();
|
||||
for bbox in &bboxes {
|
||||
assert!(bbox.x1 >= 0 && bbox.y1 >= 0, "检测框左上角不应为负");
|
||||
assert!(
|
||||
bbox.x2 <= width as i32 && bbox.y2 <= height as i32,
|
||||
"检测框右下角不应超出图片范围"
|
||||
);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
55
ddddocr-tract2/tests/ocr.rs
Normal file
55
ddddocr-tract2/tests/ocr.rs
Normal file
@@ -0,0 +1,55 @@
|
||||
//! OCR 识别与模型信息集成测试。
|
||||
|
||||
mod common;
|
||||
|
||||
use common::{model_path, sample_path};
|
||||
use ddddocr_core::traits::{Info, Loader};
|
||||
use ddddocr_core::{ModelMetadata, Normalization, Ocr, Resize};
|
||||
use ddddocr_tract2::OcrRuntime;
|
||||
use ddddocr_tract2::loader::ModelLoader;
|
||||
|
||||
/// 用官方 sml2h3 f32 模型识别验证码图片,结果不应为空。
|
||||
#[test]
|
||||
fn ocr_classification_recognizes_code_image() {
|
||||
let session = ModelLoader::default()
|
||||
.build_for_path(model_path("common_sml2h3_f32.onnx"))
|
||||
.expect("模型加载失败");
|
||||
let metadata = ModelMetadata::from_builtin_beta(
|
||||
false,
|
||||
Resize::DynamicWidth(64),
|
||||
1,
|
||||
Normalization::MinusOneToOne,
|
||||
);
|
||||
let ocr = OcrRuntime::new(session, metadata);
|
||||
|
||||
let img = image::open(sample_path("code2.png")).expect("测试图片不存在");
|
||||
let text = Ocr::builder()
|
||||
.build_with(&ocr)
|
||||
.predict(&img)
|
||||
.expect("识别过程出错")
|
||||
.into_text();
|
||||
|
||||
println!("识别结果: {text}");
|
||||
assert!(!text.is_empty(), "识别结果不应为空");
|
||||
}
|
||||
|
||||
/// 真实模型应能通过 `Info` trait 返回输入/输出张量信息。
|
||||
#[test]
|
||||
fn model_info_lists_inputs_and_outputs() -> anyhow::Result<()> {
|
||||
let session = ModelLoader::default()
|
||||
.build_for_path(model_path("common_huashi666_i64.onnx"))
|
||||
.expect("建立测试模型图失败");
|
||||
let metadata = ModelMetadata::from_builtin_beta(
|
||||
false,
|
||||
Resize::DynamicWidth(64),
|
||||
1,
|
||||
Normalization::MinusOneToOne,
|
||||
);
|
||||
let ocr = OcrRuntime::new(session, metadata);
|
||||
|
||||
let inputs = ocr.input_info()?;
|
||||
let outputs = ocr.output_info()?;
|
||||
assert!(!inputs.is_empty(), "模型应有输入张量信息");
|
||||
assert!(!outputs.is_empty(), "模型应有输出张量信息");
|
||||
Ok(())
|
||||
}
|
||||
52
ddddocr-tract2/tests/slide.rs
Normal file
52
ddddocr-tract2/tests/slide.rs
Normal file
@@ -0,0 +1,52 @@
|
||||
//! 滑块匹配集成测试。
|
||||
|
||||
mod common;
|
||||
|
||||
use common::{load_image, sample_path};
|
||||
use ddddocr_core::Slider;
|
||||
|
||||
/// 边缘模式匹配应定位到预期坐标。
|
||||
#[test]
|
||||
fn slide_match_locates_target_position() {
|
||||
let engine = Slider::new();
|
||||
let target = load_image(sample_path("target1.png")).expect("请确保 samples/target1.png 存在");
|
||||
let background = load_image(sample_path("background1.png")).expect("请确保 samples/background1.png 存在");
|
||||
|
||||
let start = std::time::Instant::now();
|
||||
let result = engine
|
||||
.slide_match(&target, &background, false)
|
||||
.expect("Slide match 执行失败");
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
println!("边缘模式匹配: {result}");
|
||||
println!("耗时: {elapsed:?}");
|
||||
|
||||
assert_eq!(result.target_x, 237);
|
||||
assert_eq!(result.target_y, 77);
|
||||
assert!(result.confidence > 0.0);
|
||||
}
|
||||
|
||||
/// 灰度对比匹配应定位到预期坐标。
|
||||
#[test]
|
||||
fn slide_comparison_locates_target_position() {
|
||||
let engine = Slider::new();
|
||||
let target = load_image(sample_path("target2.jpg")).expect("请确保 samples/target2.jpg 存在");
|
||||
let background =
|
||||
load_image(sample_path("background2.jpg")).expect("请确保 samples/background2.jpg 存在");
|
||||
|
||||
let start = std::time::Instant::now();
|
||||
let result = engine
|
||||
.slide_comparison(&target, &background)
|
||||
.expect("Slide comparison 执行失败");
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
println!(
|
||||
"灰度对比匹配: 坐标 [x: {}, y: {}], 置信度 {:.4}",
|
||||
result.target_x, result.target_y, result.confidence
|
||||
);
|
||||
println!("耗时: {elapsed:?}");
|
||||
|
||||
assert_eq!(result.target_x, 171);
|
||||
assert_eq!(result.target_y, 90);
|
||||
assert!(result.confidence > 0.0);
|
||||
}
|
||||
Reference in New Issue
Block a user