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:
2026-08-10 19:56:29 +08:00
parent fe61895926
commit 00e8ab5308
58 changed files with 2542 additions and 2237 deletions

22
ddddocr-tract2/Cargo.toml Normal file
View 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 }

View File

@@ -0,0 +1 @@
pub mod session;

View 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
View 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;

View 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;

View 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(_)));
}
}

View 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且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
Resize::Square(r1 as u32)
} else {
// 如果 word 为 false且包含 -1Python 里是高度固定为 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(_))
));
}
}

View 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());
}
}

View File

@@ -0,0 +1 @@
pub mod session;

View 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(_)));
}
}

View 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()))
}

View File

@@ -0,0 +1,7 @@
//! Tract 会话共享类型。
use std::sync::Arc;
use tract_onnx::prelude::TypedRunnableModel;
/// Tract 会话句柄:由 [`crate::loader::ModelLoader`] 构建,可直接并发执行推理。
pub type Session = Arc<TypedRunnableModel>;

View 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(())
}

View 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))
}

View 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(())
}

View 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(())
}

View 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);
}