From fe61895926424f85c4a58814d978a11c7a183ffc Mon Sep 17 00:00:00 2001 From: CNWei Date: Thu, 6 Aug 2026 19:58:54 +0800 Subject: [PATCH] =?UTF-8?q?feat(core):=20=E6=89=A9=E5=B1=95=20API=E3=80=81?= =?UTF-8?q?=E5=AE=8C=E5=96=84=E6=97=A5=E5=BF=97=E4=B8=8E=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E6=96=87=E6=A1=A3=E8=A7=84=E8=8C=83?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 公开颜色过滤与字符集限制扩展 API,修复宏路径 - 库内打印替换为 tracing 日志,清理遗留废弃代码 - 补充核心逻辑单元测试与 crate 元数据 - 开启 missing_docs 并统一 rustfmt/clippy 格式 --- ddddocr-core/Cargo.toml | 14 +- ddddocr-core/examples/quick_start.rs | 63 +++++++++ ddddocr-core/src/det/builder.rs | 1 - ddddocr-core/src/det/executor.rs | 15 +- ddddocr-core/src/error.rs | 65 ++++++++- ddddocr-core/src/lib.rs | 5 + ddddocr-core/src/ocr/builder.rs | 6 +- ddddocr-core/src/ocr/charset.rs | 42 +++++- ddddocr-core/src/ocr/color_filter.rs | 80 ++++++++++- ddddocr-core/src/ocr/executor.rs | 152 ++++----------------- ddddocr-core/src/ocr/metadata.rs | 39 ++++++ ddddocr-core/src/ocr/token_filter.rs | 89 +++++++++++- ddddocr-core/src/slide.rs | 71 +++++++--- ddddocr-core/src/traits.rs | 9 ++ ddddocr-core/src/types.rs | 11 ++ ddddocr-core/src/utils/image_convert.rs | 72 +++++++--- ddddocr-core/src/utils/image_helper.rs | 2 +- ddddocr-core/src/utils/image_processor.rs | 46 ++++++- ddddocr-core/src/utils/tensor_transform.rs | 44 +++++- ddddocr-core/tests/api_surface.rs | 72 ++++++++++ 20 files changed, 696 insertions(+), 202 deletions(-) create mode 100644 ddddocr-core/examples/quick_start.rs create mode 100644 ddddocr-core/tests/api_surface.rs diff --git a/ddddocr-core/Cargo.toml b/ddddocr-core/Cargo.toml index 8b5f74e..fbf6e49 100644 --- a/ddddocr-core/Cargo.toml +++ b/ddddocr-core/Cargo.toml @@ -3,14 +3,16 @@ name = "ddddocr-core" version = { workspace = true } edition = { workspace = true } license = { workspace = true } +description = "ddddocr-rs 的核心库:引擎无关的 OCR、目标检测与滑块匹配实现" +keywords = ["ocr", "captcha", "ddddocr", "onnx", "image"] +categories = ["multimedia::images", "computer-vision"] +# repository = "https://github.com/<用户名>/<仓库名>" # 发布前请补充 +readme = "../README.md" [dependencies] -ndarray = { workspace = true } # 继承自工作空间 - +ndarray = { workspace = true } base64 = { workspace = true } - image = { workspace = true } imageproc = { workspace = true } - -thiserror = { workspace = true } # 刚好可以开始接入你需要的标准库错误处理 -tracing={workspace = true} \ No newline at end of file +thiserror = { workspace = true } +tracing = { workspace = true } diff --git a/ddddocr-core/examples/quick_start.rs b/ddddocr-core/examples/quick_start.rs new file mode 100644 index 0000000..b047397 --- /dev/null +++ b/ddddocr-core/examples/quick_start.rs @@ -0,0 +1,63 @@ +//! 快速开始示例:演示 ddddocr-core 与引擎 crate 的解耦用法。 +//! +//! 运行:`cargo run -p ddddocr-core --example quick_start` + +use ddddocr_core::error::{Result, TensorError}; +use ddddocr_core::traits::{InferenceEngine, Info, OcrEngine}; +use ddddocr_core::types::{ModelInfo, TensorInfo}; +use ddddocr_core::{ModelMetadata, Normalization, OcrBuilder, OcrOutput, Resize}; + +/// 演示引擎:只实现接口,不接入真实 ONNX 运行时。 +struct DemoEngine { + meta: ModelMetadata, +} + +impl Info for DemoEngine { + fn input_info(&self) -> Result> { + Ok(vec![]) + } + fn output_info(&self) -> Result> { + Ok(vec![]) + } + fn model_info(&self) -> Result { + Ok(ModelInfo { + inputs: vec![], + outputs: vec![], + providers: None, + }) + } +} + +impl InferenceEngine for DemoEngine { + type Output = OcrOutput; + + fn inference(&self, input: ndarray::Array4) -> Result { + // 用全零 logits 模拟推理输出:[Steps, Classes] + let steps = input.shape()[2]; + let classes = self.meta.charset.size(); + Ok(OcrOutput::Logits(ndarray::Array2::zeros((steps, classes)))) + } +} + +impl OcrEngine for DemoEngine { + fn metadata(&self) -> &ModelMetadata { + &self.meta + } +} + +fn main() { + let engine = DemoEngine { + meta: ModelMetadata::from_static_slice( + &["", "a", "b"], + false, + Resize::Fixed(64, 64), + 1, + Normalization::ZeroToOne, + ), + }; + + let ocr = OcrBuilder::new().probability(true).build_with(&engine); + let image = image::DynamicImage::new_luma8(64, 64); + let result = ocr.predict(&image).expect("识别失败"); + println!("识别结果: {result}"); +} diff --git a/ddddocr-core/src/det/builder.rs b/ddddocr-core/src/det/builder.rs index 45f5bfe..ee298f7 100644 --- a/ddddocr-core/src/det/builder.rs +++ b/ddddocr-core/src/det/builder.rs @@ -1,7 +1,6 @@ //! 检测器构建器。 use crate::det::executor::Detector; -// use ddddocr_tract::det::session::DetSession; use crate::traits::DetEngine; /// 检测器构建器,通过 [`crate::Detector::builder`] 创建。 diff --git a/ddddocr-core/src/det/executor.rs b/ddddocr-core/src/det/executor.rs index d4bbf93..99526b0 100644 --- a/ddddocr-core/src/det/executor.rs +++ b/ddddocr-core/src/det/executor.rs @@ -11,11 +11,17 @@ use crate::{DetBuilder, DetOutput}; /// 目标检测结果:原图像素坐标系下的边界框、置信度与类别 ID。 #[derive(Debug, Clone, Copy)] pub struct DetectionResult { + /// 左上角 x 坐标。 pub x1: i32, + /// 左上角 y 坐标。 pub y1: i32, + /// 右下角 x 坐标。 pub x2: i32, + /// 右下角 y 坐标。 pub y2: i32, + /// 置信度。 pub score: f32, + /// 类别 ID。 pub class_id: u32, } @@ -36,11 +42,13 @@ pub struct Detector<'a> { } impl<'a> Detector<'a> { + /// 绑定检测引擎会话创建检测器。 pub fn new(runtime: &'a dyn DetEngine) -> Self { Detector { runtime } } + /// 创建检测器构建器。 pub fn builder() -> DetBuilder { - DetBuilder::default() + DetBuilder } } impl<'a> Detector<'a> { @@ -247,15 +255,12 @@ impl<'a> Detector<'a> { dynamic_img: &DynamicImage, ) -> Result, TensorError> { // 使用 utils crate 解码 - // let dynamic_img = image::load_from_memory(image_bytes).context("Failed to decode utils")?; let (orig_w, orig_h) = dynamic_img.dimensions(); let (input_tensor, ratio) = self.preproc(dynamic_img, (416, 416)); // tract 推理 - // let outputs = self.session.session.run(tvec!(input_tensor.into()))?; let outputs = self.runtime.inference(input_tensor)?; - // let output_array = outputs[0] // 2. 无缝、安全地解包出标准 3维 矩阵 let DetOutput::Detection(output_array) = outputs; @@ -272,9 +277,7 @@ impl<'a> Detector<'a> { expected: format!("可广播至 cls_conf 形状 {:?}", cls_conf.shape()), actual: obj_conf.shape().to_vec(), })?; - // .context("ndarray broadcasting failed for scores calculation")?; let scores = &obj_broadcast * &cls_conf; - // let scores = &pred.slice(s![.., 4..5]) * &pred.slice(s![.., 5..]); let mut boxes_xyxy = Array2::::zeros(boxes.raw_dim()); for i in 0..boxes.nrows() { diff --git a/ddddocr-core/src/error.rs b/ddddocr-core/src/error.rs index 9aaa909..06360f8 100644 --- a/ddddocr-core/src/error.rs +++ b/ddddocr-core/src/error.rs @@ -3,18 +3,20 @@ use thiserror::Error; /// 全局统一的 `Result` 别名,默认错误类型为 [`DdddError`]。 -// pub type Result = std::result::Result; pub type Result = std::result::Result; /// 顶层错误类型,聚合本库各阶段错误。 #[derive(Error, Debug)] pub enum DdddError { + /// 图像预处理阶段异常。 #[error("图像预处理失败: {0}")] Preprocess(#[from] ImagePreprocessError), + /// 推理与张量操作阶段异常。 #[error("推理与模型输入/输出张量异常: {0}")] Inference(#[from] TensorError), + /// 后处理解码阶段异常。 #[error("后处理解码错误: {0}")] Decode(#[from] DecodeError), @@ -30,46 +32,66 @@ pub enum DdddError { /// 图像预处理阶段错误类型。 #[derive(Error, Debug)] pub enum ImagePreprocessError { + /// ndarray 基础操作失败。 #[error("图片转矩阵(ndarray)基础操作失败: {0}")] Ndarray(#[from] ndarray::ShapeError), + /// 图像矩阵维度不合规。 #[error("图像矩阵维度不合规!预期: {expected},实际图像形状: {actual:?}")] InvalidDimensions { + /// 期望的维度描述。 expected: String, + /// 实际的图像形状。 actual: Vec, }, + /// 图像缓冲区长度与分辨率/通道数不匹配。 #[error( "图像缓冲区长度不匹配!预期大小: {expected},实际大小: {actual} (分辨率: {width}x{height}, 通道数: {channels})" )] BufferLengthMismatch { + /// 期望的缓冲区长度。 expected: usize, + /// 实际的缓冲区长度。 actual: usize, + /// 图像宽度。 width: u32, + /// 图像高度。 height: u32, + /// 图像通道数。 channels: usize, }, + /// 不支持的图像通道数。 #[error("不支持的图像通道数: {0} (仅支持单通道灰度L、3通道RGB、4通道RGBA)")] UnsupportedChannels(usize), + /// HSV 颜色区间参数非法。 #[error("HSV 颜色区间参数非法: {0}")] InvalidHsvRange(String), + /// 未知的颜色预设名称。 #[error("不支持的颜色预设名称: {0}")] UnknownColorPreset(String), + /// 颜色过滤器配置无效。 #[error("颜色过滤器配置无效或初始化失败: {0}")] FilterConfigInvalid(String), + /// 图像维度不匹配。 #[error("图像维度不匹配!{0}")] MismatchDimensions(String), + /// 滑块模板尺寸大于背景图。 #[error("滑块模板尺寸 [{target_w}x{target_h}] 大于背景图 [{bg_w}x{bg_h}]")] TargetExceedsBackground { + /// 滑块模板宽度。 target_w: usize, + /// 滑块模板高度。 target_h: usize, + /// 背景图宽度。 bg_w: usize, + /// 背景图高度。 bg_h: usize, }, } @@ -77,21 +99,28 @@ pub enum ImagePreprocessError { /// 推理与张量操作阶段错误类型。 #[derive(Error, Debug)] pub enum TensorError { + /// 推理引擎内部异常。 #[error("推理引擎内部发生异常: {0}")] Engine(String), + /// 模型张量维度不匹配。 #[error("模型张量维度不匹配!预期: {expected},实际 Tensor 形状: {actual:?}")] DimensionMismatch { + /// 期望的维度描述。 expected: String, + /// 实际的 Tensor 形状。 actual: Vec, }, + /// OCR Logits 矩阵变形失败。 #[error("OCR Logits 矩阵变形失败: {0}")] LogitsDimensionMismatch(#[from] ndarray::ShapeError), + /// 张量内存不连续。 #[error("内存不连续,无法执行零拷贝操作")] NonContiguousMemory, + /// 未知的模型输出格式。 #[error("未知的模型输出格式")] UnknownOutputFormat, } @@ -99,6 +128,7 @@ pub enum TensorError { /// 算法解码阶段错误类型。 #[derive(Error, Debug)] pub enum DecodeError { + /// CTC 解码异常。 #[error("CTC 解码异常: {0}")] Ctc(String), } @@ -112,6 +142,7 @@ impl DdddError { DdddError::Other(error.into()) } + /// 是否为图片维度不合规错误。 pub fn is_invalid_dimensions(&self) -> bool { matches!( self, @@ -119,6 +150,7 @@ impl DdddError { ) } + /// 是否因通道数不合规而失败。 pub fn is_unsupported_channels(&self) -> bool { matches!( self, @@ -126,3 +158,34 @@ impl DdddError { ) } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn new_wraps_third_party_error() { + let io_err = std::io::Error::new(std::io::ErrorKind::Other, "boom"); + let err = DdddError::new(io_err); + assert!(matches!(err, DdddError::Other(_))); + } + + #[test] + fn preprocess_conversion_and_predicates() { + let e: DdddError = ImagePreprocessError::UnsupportedChannels(2).into(); + assert!(e.is_unsupported_channels()); + + let e2: DdddError = ImagePreprocessError::InvalidDimensions { + expected: "x".into(), + actual: vec![0], + } + .into(); + assert!(e2.is_invalid_dimensions()); + } + + #[test] + fn decode_conversion() { + let e: DdddError = DecodeError::Ctc("bad".into()).into(); + assert!(matches!(e, DdddError::Decode(_))); + } +} diff --git a/ddddocr-core/src/lib.rs b/ddddocr-core/src/lib.rs index 136266b..947d4a3 100644 --- a/ddddocr-core/src/lib.rs +++ b/ddddocr-core/src/lib.rs @@ -6,6 +6,8 @@ //! //! 完整可运行示例见 `ddddocr-core/examples/quick_start.rs`。 +#![warn(missing_docs)] + mod det; /// 分层错误类型。 pub mod error; @@ -28,11 +30,14 @@ pub use crate::slide::{SlideResult, Slider}; /// OCR 模型的统一输出枚举,由推理引擎产出,供 [`Ocr`] 后处理。 pub enum OcrOutput { + /// 索引序列输出(CTC 解码输入)。 Indices(ndarray::Array1), + /// Logits 矩阵输出 `[Steps, Classes]`。 Logits(ndarray::Array2), } /// 目标检测模型的统一输出枚举,由推理引擎产出,供 [`Detector`] 后处理。 pub enum DetOutput { + /// 原始检测输出张量。 Detection(ndarray::Array3), } diff --git a/ddddocr-core/src/ocr/builder.rs b/ddddocr-core/src/ocr/builder.rs index b1df876..c640c55 100644 --- a/ddddocr-core/src/ocr/builder.rs +++ b/ddddocr-core/src/ocr/builder.rs @@ -20,6 +20,7 @@ pub struct OcrBuilder { impl OcrBuilder { // 初始化任务,设置默认参数 + /// 创建默认配置的构建器。 pub fn new() -> Self { Self { png_fix: false, // 默认值 @@ -28,15 +29,18 @@ impl OcrBuilder { charset_restrict: None, } } + /// 设置是否修复 PNG 透明背景问题。 pub fn png_fix(mut self, value: bool) -> Self { self.png_fix = value; self } + /// 设置是否返回概率信息。 pub fn probability(mut self, value: bool) -> Self { self.probability = value; self } + /// 设置颜色过滤约束。 pub fn color_filter(mut self, filter: T) -> Self where T: ColorFilter + Send + Sync + 'static, @@ -45,6 +49,7 @@ impl OcrBuilder { self } + /// 设置字符集限制。 pub fn charset_restrict(mut self, restrict: T) -> Self where T: TokenFilter + Send + Sync + 'static, @@ -66,7 +71,6 @@ impl OcrBuilder { None => None, }; - // Ocr::new(session, self) Ocr { runtime, png_fix: self.png_fix, // 原地解构出来 diff --git a/ddddocr-core/src/ocr/charset.rs b/ddddocr-core/src/ocr/charset.rs index a128fc5..d60c8c3 100644 --- a/ddddocr-core/src/ocr/charset.rs +++ b/ddddocr-core/src/ocr/charset.rs @@ -9,22 +9,22 @@ use std::collections::HashMap; /// 字符集:token 列表与索引的双向映射。 #[derive(Debug, Clone)] pub struct Charset { + /// 字符集 token 列表。 // 使用 Cow 统一静态切片和动态读取的 Vec,内部实现真正的零拷贝 pub tokens: Vec>, + /// 字符到索引的反查表。 // 反向查找表,保证字符转索引为 O(1) pub char_to_idx: HashMap, usize>, - // 当前处于激活状态的有效索引缓存 (用于 CTC 解码前的过滤加速) - // pub valid_indices: HashSet, } impl Charset { // 内部底层统一收拢构造 + /// 从 token 列表构建字符集。 pub fn new(tokens: Vec>) -> Self { let mut char_to_idx = HashMap::with_capacity(tokens.len()); for (idx, token) in tokens.iter().enumerate() { char_to_idx.entry(token.clone()).or_insert(idx); // 如果字符集有重复,保留第一个遇到的索引 (符合 Python .index 逻辑) - // char_to_idx.entry(token.to_string()).or_insert(idx); } Self { @@ -49,9 +49,11 @@ impl Charset { self.tokens.get(index).map(|cow| cow.as_ref()) } + /// 判断字符是否在字符集中。 pub fn is_valid_char(&self, char_str: &str) -> bool { - self.char_to_idx.get(char_str).is_some() + self.char_to_idx.contains_key(char_str) } + /// 返回字符集大小。 pub fn size(&self) -> usize { self.tokens.len() } @@ -66,4 +68,36 @@ impl std::fmt::Display for Charset { } } +#[cfg(test)] +mod tests { + use super::*; + fn sample_tokens() -> Vec> { + vec![Cow::Borrowed(""), Cow::Borrowed("a"), Cow::Borrowed("b")] + } + + #[test] + fn char_to_index_roundtrip() { + let cs = Charset::new(sample_tokens()); + assert_eq!(cs.char_to_index("a"), 1); + assert_eq!(cs.char_to_index("z"), -1); + assert_eq!(cs.index_to_char_ref(2), Some("b")); + assert_eq!(cs.index_to_char_ref(99), None); + } + + #[test] + fn duplicate_tokens_keep_first_index() { + let cs = Charset::new(vec![Cow::Borrowed("x"), Cow::Borrowed("x")]); + assert_eq!(cs.char_to_index("x"), 0); + assert_eq!(cs.size(), 2); + } + + #[test] + fn is_valid_char_and_size() { + let cs = Charset::new(sample_tokens()); + assert!(cs.is_valid_char("")); + assert!(cs.is_valid_char("a")); + assert!(!cs.is_valid_char("A")); + assert_eq!(cs.size(), 3); + } +} diff --git a/ddddocr-core/src/ocr/color_filter.rs b/ddddocr-core/src/ocr/color_filter.rs index f2f4ebd..3a4e3f3 100644 --- a/ddddocr-core/src/ocr/color_filter.rs +++ b/ddddocr-core/src/ocr/color_filter.rs @@ -48,8 +48,7 @@ pub fn apply_to_image( // 3. 将扁平字节数组重新打包回 DynamicImage 容器 let filtered_buffer = ImageBuffer::, Vec>::from_raw(width, height, raw_pixels) - // .ok_or_else(|| anyhow!("图像缓冲重新组装失败,维度与数据大小不匹配"))?; - .ok_or_else(|| ImagePreprocessError::BufferLengthMismatch { + .ok_or(ImagePreprocessError::BufferLengthMismatch { expected: expected_len, actual: actual_len, width, @@ -62,11 +61,14 @@ pub fn apply_to_image( #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] /// HSV 颜色区间,下界与上界各为 `(H, S, V)`。 pub struct HsvRange { + /// 区间下界 `(H, S, V)`。 pub lower: (u8, u8, u8), // (H, S, V) + /// 区间上界 `(H, S, V)`。 pub upper: (u8, u8, u8), // (H, S, V) } impl HsvRange { + /// 创建 HSV 区间。 pub const fn new(lower: (u8, u8, u8), upper: (u8, u8, u8)) -> Self { Self { lower, upper } } @@ -76,7 +78,6 @@ impl HsvRange { pub fn validate(&self) -> Result<(), ImagePreprocessError> { // 1. 校验 H 通道边界 (OpenCV 中 H 范围是 0-180) if self.lower.0 > 180 || self.upper.0 > 180 { - // return Err("H通道值必须在 0-180 范围内".to_string()); return Err(ImagePreprocessError::InvalidHsvRange( "H通道值必须在 0-180 范围内".to_string(), )); @@ -85,7 +86,6 @@ impl HsvRange { // 2. 校验下界不能大于上界 if self.lower.0 > self.upper.0 || self.lower.1 > self.upper.1 || self.lower.2 > self.upper.2 { - // return Err("HSV范围下界不能大于上界".to_string()); return Err(ImagePreprocessError::InvalidHsvRange( "HSV范围下界不能大于上界".to_string(), )); @@ -97,16 +97,27 @@ impl HsvRange { #[derive(Debug, Clone, PartialEq, Eq)] /// 颜色预设:常见颜色对应的 HSV 区间集合。 pub enum ColorPreset { + /// 红色。 Red, + /// 蓝色。 Blue, + /// 绿色。 Green, + /// 黄色。 Yellow, + /// 橙色。 Orange, + /// 紫色。 Purple, + /// 青色。 Cyan, + /// 黑色。 Black, + /// 白色。 White, + /// 灰色。 Gray, + /// 自定义区间列表。 Custom(Vec), } @@ -193,7 +204,6 @@ impl FromStr for ColorPreset { "black" => Ok(ColorPreset::Black), "white" => Ok(ColorPreset::White), "gray" => Ok(ColorPreset::Gray), - // _ => Err(format!("不支持的颜色预设: {}", s)), _ => Err(ImagePreprocessError::UnknownColorPreset(s.to_string())), } } @@ -205,12 +215,15 @@ impl FromStr for ColorPreset { /// 颜色匹配上下文:当前像素的 HSV 值。 pub struct PixelCtx { + /// 当前像素的 HSV 值。 pub hsv: (u8, u8, u8), } /// 颜色过滤约束接口:提供一组 HSV 区间。 pub trait ColorFilter { + /// 将有效区间追加到目标容器。 fn append_ranges(&self, target: &mut Vec); + /// 预估有效区间数量。 fn estimated_count(&self) -> usize; /// 验证过滤器配置是否合法,默认直接放行。 fn validate_self(&self) -> Result<(), ImagePreprocessError> { @@ -258,6 +271,7 @@ impl ColorFilter for ColorPreset { /// 多路颜色“或”逻辑组合子(并集网络) pub struct MultiOrColorRestrict<'a> { + /// 参与「或」组合的过滤器列表。 pub filters: Vec<&'a dyn ColorFilter>, } @@ -298,3 +312,59 @@ macro_rules! color_any_of { } }; } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn hsv_range_validate() { + assert!(HsvRange::new((0, 0, 0), (180, 255, 255)).validate().is_ok()); + assert!( + HsvRange::new((181, 0, 0), (255, 255, 255)) + .validate() + .is_err() + ); + assert!( + HsvRange::new((20, 50, 50), (10, 255, 255)) + .validate() + .is_err() + ); + } + + #[test] + fn color_preset_matches_counts() { + assert_eq!(ColorPreset::Red.matches().len(), 2); + assert_eq!(ColorPreset::Blue.matches().len(), 1); + let custom = ColorPreset::Custom(vec![HsvRange::new((0, 0, 0), (1, 1, 1))]); + assert_eq!(custom.matches().len(), 1); + } + + #[test] + fn color_preset_from_str() { + assert_eq!("red".parse::().unwrap(), ColorPreset::Red); + assert!("pink".parse::().is_err()); + } + + #[test] + fn color_any_of_collects() { + let ranges = crate::color_any_of!(ColorPreset::Red, ColorPreset::Blue) + .collect_to_vec() + .unwrap() + .unwrap(); + assert_eq!(ranges.len(), 3); + } + + #[test] + fn apply_to_image_whitens_non_matching() { + let img = DynamicImage::ImageRgb8(image::ImageBuffer::from_pixel( + 2, + 2, + image::Rgb([255, 0, 0]), + )); + let filtered = apply_to_image(&img, ColorPreset::Blue.matches()).unwrap(); + for p in filtered.to_rgb8().pixels() { + assert_eq!(p.0, [255, 255, 255]); + } + } +} diff --git a/ddddocr-core/src/ocr/executor.rs b/ddddocr-core/src/ocr/executor.rs index 1f62c10..c238dd3 100644 --- a/ddddocr-core/src/ocr/executor.rs +++ b/ddddocr-core/src/ocr/executor.rs @@ -2,9 +2,8 @@ use crate::ocr::metadata::Resize; -use crate::ocr::color_filter::{HsvRange, apply_to_image}; -// use ddddocr_tract::session::{ModelOutput, OcrSession}; use crate::error::{ImagePreprocessError, Result, TensorError}; +use crate::ocr::color_filter::{HsvRange, apply_to_image}; use crate::traits::OcrEngine; use crate::utils::image_convert::png_rgba_white_preprocess; use crate::utils::image_processor::{convert_to_grayscale, resize_image}; @@ -13,7 +12,7 @@ use image::DynamicImage; use ndarray::ArrayView2; use std::borrow::Cow; use std::fmt; -use tracing::warn; +use tracing::{debug, warn}; /// OCR 识别结果:纯文本或携带概率的文本。 #[derive(Debug, Clone)] pub enum OcrResult { @@ -21,6 +20,7 @@ pub enum OcrResult { Text(String), /// 携带概率的结果(`probability = true` 时返回)。 Probability { + /// 识别出的文本。 text: String, /// 全量概率矩阵 `[Steps, Classes]`。 probabilities: Vec>, @@ -28,7 +28,10 @@ pub enum OcrResult { confidence: f64, }, /// 不支持的模型或未知输出。 - Unsupported { message: String }, + Unsupported { + /// 不支持原因说明。 + message: String, + }, } impl OcrResult { /// 消费自身并提取最终文本。 @@ -114,6 +117,7 @@ pub struct Ocr<'a> { impl<'a> Ocr<'a> { // 初始化任务,设置默认参数 + /// 绑定引擎会话创建 OCR 识别器。 pub fn new(runtime: &'a dyn OcrEngine) -> Self { Ocr { runtime, @@ -123,6 +127,7 @@ impl<'a> Ocr<'a> { final_charset_indices: None, } } + /// 创建 OCR 构建器。 pub fn builder() -> OcrBuilder { OcrBuilder::default() } @@ -130,7 +135,7 @@ impl<'a> Ocr<'a> { impl<'a> Ocr<'a> { /// 对输入图像执行 OCR 识别并返回结果。 pub fn predict(&self, image: &DynamicImage) -> Result { - println!("当前颜色过滤器状态: {:?}", self.final_color_ranges); + debug!("当前颜色过滤器状态: {:?}", self.final_color_ranges); // ===================================================================== // 管道节点 1: 颜色过滤流水线 @@ -139,10 +144,6 @@ impl<'a> Ocr<'a> { // ===================================================================== let img_cow = match &self.final_color_ranges { Err(err_msg) => { - // return Err(anyhow::anyhow!( - // "颜色过滤器初始化失败,全链路短路: {}", - // err_msg - // )); return Err(ImagePreprocessError::FilterConfigInvalid( err_msg.to_string(), ))?; @@ -162,17 +163,6 @@ impl<'a> Ocr<'a> { let raw_tensor = self.runtime.inference(tensor)?; // 3. 后处理分流:直接返回 OcrResult - // let ocr_output = match raw_tensor.datum_type() { - // DatumType::I64 => self.process_i64_tensor(raw_tensor)?, - // DatumType::F32 => self.process_f32_tensor(raw_tensor)?, - // _ => OcrResult::Unsupported { - // message: format!("不支持的模型输出数据类型: {:?}", raw_tensor.datum_type()), - // }, - // }; - - // let raw_indices = self.ocr.extract_indices_from_tensor(&raw_tensor)?; - // // 步骤 2: 将索引切片 `&[i64]` 传给解码器进行 CTC 去重和字符映射 - // let final_text = self.ctc_decode_to_string(&raw_indices); let ocr_output = self.process_model_output(raw_tensor)?; Ok(ocr_output) } @@ -216,7 +206,7 @@ impl<'a> Ocr<'a> { 1 => { let gray_img = convert_to_grayscale(&resized_img); - let array = ndarray::Array4::from_shape_fn( + ndarray::Array4::from_shape_fn( (1, 1, target_h as usize, target_w as usize), |(_, _, y, x)| { let pixel = gray_img.get_pixel(x as u32, y as u32)[0] as f32; @@ -224,15 +214,14 @@ impl<'a> Ocr<'a> { // (pixel / 255.0 - 0.5) / 0.5 norm.normalize(pixel) }, - ); - array + ) } // --- 情况 B: 三通道(RGB),对应 Python 的 transpose(2, 0, 1) 的 CHW 布局 --- 3 => { let rgb_img = resized_img.to_rgb8(); - let array = ndarray::Array4::from_shape_fn( + ndarray::Array4::from_shape_fn( (1, 3, target_h as usize, target_w as usize), |(_, c, y, x)| { let pixel = rgb_img.get_pixel(x as u32, y as u32)[c] as f32; @@ -240,9 +229,7 @@ impl<'a> Ocr<'a> { // (pixel / 255.0 - 0.5) / 0.5 norm.normalize(pixel) }, - ); - // Tensor::from(array) - array + ) } // _ => return Err(anyhow::anyhow!("不支持的通道数配置: {}", meta.channel)), @@ -255,15 +242,11 @@ impl<'a> Ocr<'a> { Ok(array4) } - // 这段代码未来直接放入 ddddocr-core fn process_model_output(&self, output: OcrOutput) -> Result { match output { OcrOutput::Indices(array1) => { - // 对应你原来的 process_i64_tensor - let slice = array1 - .as_slice() - // .ok_or_else(|| anyhow::anyhow!("内存不连续,无法执行零拷贝解码"))?; - .ok_or_else(|| TensorError::NonContiguousMemory)?; + // 对应原来的 process_i64_tensor + let slice = array1.as_slice().ok_or(TensorError::NonContiguousMemory)?; let final_text = self.ctc_decode_to_string(slice); if self.probability { @@ -277,7 +260,7 @@ impl<'a> Ocr<'a> { } } OcrOutput::Logits(matrix_view) => { - // 对应你原来的 process_f32_tensor + // 对应原来的 process_f32_tensor // 注意:此时的 matrix_view 已经是干净的标准的 ndarray::Array2,且保证是 [Steps, Classes] 2D 形状 if self.probability { let (probabilities_list, confidence, predicted_indices) = @@ -332,6 +315,7 @@ impl<'a> Ocr<'a> { None => tokens.iter().map(|cow| cow.as_ref()).collect(), } } + /// 返回当前生效的可用 token 数量。 pub fn valid_size(&self) -> usize { match &self.final_charset_indices { Some(indices) => indices.len(), @@ -390,100 +374,10 @@ impl<'a> Ocr<'a> { (probabilities_list, confidence, predicted_indices) } - /// 变体 A 专属提取器:直接从 I64 Tensor 零拷贝提取 CTC 文本与初始概率包 - // fn process_i64_tensor(&self, raw_tensor: Tensor) -> anyhow::Result { - // // 1. 拿到底层的动态维度只读视图 - // let view = raw_tensor.to_array_view::()?; - // - // // 2. 索要底层连续的只读切片引用 - // let slice = view - // .as_slice() - // .ok_or_else(|| anyhow::anyhow!("I64 模型输出内存不连续,无法执行零拷贝解码"))?; - // - // // 3. 直接喂给 CTC 解码器(无任何物理克隆开销) - // let final_text = self.ctc_decode_to_string(slice); - // - // // 4. 组装返回 - // if self.probability { - // Ok(OcrResult::Probability { - // text: final_text, - // probabilities: vec![], // I64 模型物理上丢失了全量 Logits 分值网,降级处理 - // confidence: 1.0, // 判定即百分之百置信 - // }) - // } else { - // Ok(OcrResult::Text(final_text)) - // } - // } - // /// 变体二(F32)的总体管线:负责降维,并分流文本和概率 - // fn process_f32_tensor(&self, raw_tensor: Tensor) -> anyhow::Result { - // let shape = raw_tensor.shape(); - // println!("模型输出shape数据: {:?}", shape); - // let view = raw_tensor.to_array_view::()?; - // - // // 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗 - // let (steps, classes, data_dyn_view) = match shape.len() { - // 3 => { - // if shape[1] == 1 { - // // 形状: [Steps, 1, Classes] -> 你的原有逻辑 - // (shape[0], shape[2], view.into_dyn()) - // } else if shape[0] == 1 { - // // 形状: [1, Steps, Classes] -> 另一种常见导出格式 - // (shape[1], shape[2], view.into_dyn()) - // } else { - // // 默认取第一个 batch: [Batch, Steps, Classes] - // // 使用 slice 对应 Python 的 output[0, :, :] - // let sliced = view.slice(s![0, .., ..]); - // (shape[1], shape[2], sliced.into_dyn()) - // } - // } - // // 形状: [Steps, Classes] -> 已经剥离了 Batch 维度 - // 2 => (shape[0], shape[1], view.into_dyn()), - // // 形状: [Classes] -> 单字符输出(对应 Python 的 ndim == 0 保护逻辑) - // // 我们把它虚构成一个 [1, Classes] 的 2D 矩阵来复用后面的 argmax 逻辑 - // 1 => (1, shape[0], view.into_dyn()), - // _ => return Err(anyhow::anyhow!("不支持的输出维度: {:?}", shape)), - // }; - // let matrix_cow = data_dyn_view - // .to_shape(Ix2(steps, classes)) - // .map_err(|e| anyhow::anyhow!("转换为2D静态矩阵失败: {:?}", e))?; - // - // let matrix_view: ArrayView2 = matrix_cow.view(); - // - // // 2. 根据业务参数明确分流 - // if self.probability { - // // 走向 B1:调用刚刚拆分出来的“全量概率计算器” - // let (probabilities_list, confidence, predicted_indices) = - // self.compute_f32_full_probability(matrix_view); - // // 5. 执行 CTC 解码 - // let final_text = self.ctc_decode_to_string(&predicted_indices); - // - // Ok(OcrResult::Probability { - // text: final_text, - // probabilities: probabilities_list, - // confidence: confidence as f64, - // }) - // } else { - // // 走向 B2:极速免 Softmax 提取纯文本(代码保持原地提取,简单短小不需要再拆) - // let predicted_indices: Vec = matrix_view - // .outer_iter() - // .map(|row| { - // row.iter() - // .enumerate() - // .max_by(|(_, a), (_, b)| a.total_cmp(b)) - // .map(|(idx, _)| idx as i64) - // .unwrap_or(0) - // }) - // .collect(); - // - // let final_text = self.ctc_decode_to_string(&predicted_indices); - // Ok(OcrResult::Text(final_text)) - // } - // } fn ctc_decode_to_string(&self, predicted_indices: &[i64]) -> String { - println!("indices模型输出原始数据: {:?}", predicted_indices); + debug!("indices模型输出原始数据: {:?}", predicted_indices); let charset = &self.runtime.metadata().charset; let tokens = &charset.tokens; - // let valid_indices = &charset.valid_indices; // 对应 _ctc_decode_indices 的逻辑:去重、去 blank (0) let mut res = String::new(); @@ -509,10 +403,10 @@ impl<'a> Ocr<'a> { // 史诗级加速点:如果是 None,说明没限制,根本不进入分支,直接放行! // 只有当有具体限制(Some)时,才去跑 4-5 次 CPU 寄存器级别的二分查找 - if let Some(ref indices) = self.final_charset_indices { - if indices.binary_search(&u_idx).is_err() { - continue; - } + if let Some(ref indices) = self.final_charset_indices + && indices.binary_search(&u_idx).is_err() + { + continue; } // 5. 字符映射 diff --git a/ddddocr-core/src/ocr/metadata.rs b/ddddocr-core/src/ocr/metadata.rs index 2d00e66..aad5d71 100644 --- a/ddddocr-core/src/ocr/metadata.rs +++ b/ddddocr-core/src/ocr/metadata.rs @@ -17,6 +17,7 @@ pub enum Normalization { } impl Normalization { + /// 对像素值执行归一化。 #[inline(always)] pub fn normalize(&self, pixel: f32) -> f32 { match self { @@ -40,16 +41,20 @@ pub enum Resize { /// OCR 模型元数据:字符集、缩放策略、通道数与归一化配置。 #[derive(Debug, Clone)] pub struct ModelMetadata { + /// 字符集。 pub charset: Charset, /// 是否为单字识别模型 pub word: bool, + /// 缩放策略。 pub resize: Resize, /// 图像通道数(1 或 3) pub channel: u8, + /// 像素归一化配置。 pub normalization: Normalization, } impl ModelMetadata { + /// 创建模型元数据。 pub fn new( charset: Charset, word: bool, @@ -84,3 +89,37 @@ impl ModelMetadata { } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn normalization_zero_to_one() { + let n = Normalization::ZeroToOne; + assert_eq!(n.normalize(0.0), 0.0); + assert_eq!(n.normalize(255.0), 1.0); + } + + #[test] + fn normalization_minus_one_to_one() { + let n = Normalization::MinusOneToOne; + assert_eq!(n.normalize(0.0), -1.0); + assert_eq!(n.normalize(255.0), 1.0); + } + + #[test] + fn from_static_slice_builds_charset() { + let meta = ModelMetadata::from_static_slice( + &["", "a"], + false, + Resize::Fixed(64, 64), + 1, + Normalization::ZeroToOne, + ); + assert_eq!(meta.charset.size(), 2); + assert_eq!(meta.charset.char_to_index("a"), 1); + assert_eq!(meta.channel, 1); + assert!(!meta.word); + } +} diff --git a/ddddocr-core/src/ocr/token_filter.rs b/ddddocr-core/src/ocr/token_filter.rs index be00fb2..b5aa91e 100644 --- a/ddddocr-core/src/ocr/token_filter.rs +++ b/ddddocr-core/src/ocr/token_filter.rs @@ -1,21 +1,25 @@ //! 字符集限制:按字符属性或索引过滤识别范围。 use std::borrow::Cow; +use tracing::warn; /// 字符集校验上下文:当前 token 的文本与索引。 pub struct ValidationCtx<'a> { - pub text: &'a str, // 当前 Token 的文本内容 + /// 当前 token 的文本内容。 + pub text: &'a str, // 当前 Token 的文本内容 + /// 当前 token 的 ID 索引。 pub token_id: usize, // 当前 Token 的 ID 索引 } /// 字符集限制接口:决定某个 token 是否放行。 pub trait TokenFilter { + /// 判断 token 是否放行。 fn matches(&self, ctx: &ValidationCtx) -> bool; /// 预估匹配数量的容量提示。 fn estimated_capacity(&self) -> usize { 128 } - + /// 遍历全量字符集筛选可用索引(放行 CTC blank、排序去重、空交集返回 `None`)。 fn apply_to_charset(&self, tokens: &[Cow]) -> Option> { let mut has_any_match = false; @@ -49,7 +53,7 @@ pub trait TokenFilter { // 3. 终极防御:如果整个模型字符集除了 Blank,一个都没对上,直接退化为 None(全量识别) if !has_any_match { - println!("警告:当前限制策略与模型字符集完全没有交集!已自动恢复全量识别。"); + warn!("当前限制策略与模型字符集完全没有交集,已自动恢复全量识别"); None } else { // 4. 排序并去重,为 Ocr 引擎后续进行极其高频的『二分查找』筑起绝对安全的底层保障 @@ -63,9 +67,13 @@ pub trait TokenFilter { /// 按字符属性限制:数字、大小写字母或自定义列表。 #[derive(Debug, Clone, PartialEq, Eq)] pub enum CharRestrict { + /// 仅数字。 Digit, + /// 仅小写字母。 Lowercase, + /// 仅大写字母。 Uppercase, + /// 自定义字符列表。 CustomList(Vec), } @@ -90,8 +98,11 @@ impl TokenFilter for CharRestrict { /// 按索引限制:前 N 个、索引范围或索引列表。 #[derive(Debug, Clone, PartialEq, Eq)] pub enum IdRestrict { + /// 前 N 个索引。 TopN(usize), + /// 指定索引范围。 IdRange(std::ops::Range), + /// 指定索引列表。 IdList(Vec), } @@ -117,6 +128,7 @@ impl TokenFilter for IdRestrict { /// 多路“或”逻辑组合子(支持 N 个规则无缝并集) pub struct MultiOrRestrict<'a> { + /// 参与「或」组合的过滤器列表。 pub filters: Vec<&'a dyn TokenFilter>, } @@ -147,3 +159,74 @@ macro_rules! any_of { } }; } + +#[cfg(test)] +mod tests { + use super::*; + + fn tokens() -> Vec> { + vec![ + Cow::Borrowed(""), + Cow::Borrowed("1"), + Cow::Borrowed("a"), + Cow::Borrowed("A"), + Cow::Borrowed("!"), + ] + } + + #[test] + fn char_restrict_digit() { + let indices = CharRestrict::Digit.apply_to_charset(&tokens()).unwrap(); + assert_eq!(indices, vec![0, 1]); + } + + #[test] + fn char_restrict_case() { + assert_eq!( + CharRestrict::Lowercase.apply_to_charset(&tokens()).unwrap(), + vec![0, 2] + ); + assert_eq!( + CharRestrict::Uppercase.apply_to_charset(&tokens()).unwrap(), + vec![0, 3] + ); + } + + #[test] + fn custom_list() { + let r = CharRestrict::CustomList(vec!["!".into(), "a".into()]); + assert_eq!(r.apply_to_charset(&tokens()).unwrap(), vec![0, 2, 4]); + } + + #[test] + fn no_intersection_falls_back_to_none() { + let r = CharRestrict::CustomList(vec!["z".into()]); + assert_eq!(r.apply_to_charset(&tokens()), None); + } + + #[test] + fn id_restrict_variants() { + assert_eq!( + IdRestrict::TopN(3).apply_to_charset(&tokens()).unwrap(), + vec![0, 1, 2] + ); + assert_eq!( + IdRestrict::IdRange(1..3) + .apply_to_charset(&tokens()) + .unwrap(), + vec![0, 1, 2] + ); + assert_eq!( + IdRestrict::IdList(vec![0, 4]) + .apply_to_charset(&tokens()) + .unwrap(), + vec![0, 4] + ); + } + + #[test] + fn multi_or_restrict_unions() { + let combined = crate::any_of!(CharRestrict::Digit, CharRestrict::Uppercase); + assert_eq!(combined.apply_to_charset(&tokens()).unwrap(), vec![0, 1, 3]); + } +} diff --git a/ddddocr-core/src/slide.rs b/ddddocr-core/src/slide.rs index b1af8fa..3e9b5cb 100644 --- a/ddddocr-core/src/slide.rs +++ b/ddddocr-core/src/slide.rs @@ -14,14 +14,18 @@ use imageproc::region_labelling::{Connectivity, connected_components}; use imageproc::template_matching::{MatchTemplateMethod, match_template}; use ndarray::{ArrayView2, ArrayView3}; use std::fmt; +use tracing::debug; /// 滑块匹配结果:检测中心坐标与置信度。 #[derive(Debug)] pub struct SlideResult { /// 检测中心坐标 `[x, y]`。 pub target: [i32; 2], + /// 检测中心的 x 坐标。 pub target_x: i32, + /// 检测中心的 y 坐标。 pub target_y: i32, + /// 匹配置信度。 pub confidence: f64, } impl fmt::Display for SlideResult { @@ -35,9 +39,11 @@ impl fmt::Display for SlideResult { } /// 滑块匹配服务:提供模板匹配与差异比较两种识别模式。 +#[derive(Default)] pub struct Slider; impl Slider { + /// 创建滑块匹配服务。 pub fn new() -> Self { Self } @@ -113,7 +119,6 @@ impl Slider { let background_label = Luma([0u8]); let labelled = connected_components(&cleaned, Connectivity::Eight, background_label); - // // 统计每个标签出现的频率(即面积) // 4. 寻找最大连通区域 (对应 findContours + max area) if let Some(max_label) = image_processor::find_contours_and_max(&labelled) { // 5. 计算最大区域的边界框 (对应 cv2.boundingRect) @@ -139,7 +144,6 @@ impl Slider { } /// 模板匹配核心实现(对应 Python 的 `_perform_slide_match`)。 - // 在 SlideEngine 中修改此入口进行测试 fn perform_slide_match( &self, target: ArrayView3, @@ -158,7 +162,6 @@ impl Slider { } if th > bh || tw > bw { return Err(ImagePreprocessError::TargetExceedsBackground { - // "尺寸不匹配:滑块模板(target)尺寸 [{}x{}] 不能大于背景图(background) [{}x{}]", target_w: tw, target_h: th, bg_w: bw, @@ -195,7 +198,6 @@ impl Slider { // 转换逻辑 (假设你已经有方法转回 ImageBuffer) let t_buf = ndarray_to_luma8(target); let b_buf = ndarray_to_luma8(background); - // t_buf.save("debug_rust_target.png").unwrap(); // 2. 调用 imageproc 的 NCC 算法 (等价于 cv2.TM_CCOEFF_NORMED) // 模板匹配 (完全对齐 cv2.matchTemplate(..., cv2.TM_CCOEFF_NORMED)) @@ -204,18 +206,13 @@ impl Slider { &t_buf, MatchTemplateMethod::CrossCorrelationNormalized, ); - // save_rust_result(&result, "debug_rust_target2.png"); // 3. 寻找最大值 (等价于 cv2.minMaxLoc) let (max_val, max_loc) = min_max_loc(&result); // 4. 计算中心点 (与 Python 逻辑完全一致) let (th, tw) = target.dim(); - let (center_x, center_y) = - image_processor::calculate_center(max_loc, tw as usize, th as usize); - // println!("Rust Target Width (tw): {}", tw); - // println!("Rust Best Max Loc X: {}", max_loc.0); - // println!("Rust Final Center X: {}", center_x); + let (center_x, center_y) = image_processor::calculate_center(max_loc, tw, th); SlideResult { target: [center_x, center_y], target_x: center_x, @@ -240,9 +237,6 @@ impl Slider { let target_edges = canny(&t_buf, 50.0, 150.0); let background_edges = canny(&b_buf, 50.0, 150.0); - // target_edges.save("debug_target_edges.png").ok(); - // background_edges.save("debug_bg_edges.png").ok(); - // 3. 模板匹配 (完全对齐 cv2.matchTemplate(..., cv2.TM_CCOEFF_NORMED)) // 在边缘图上计算归一化互相关系数 let result = match_template( @@ -256,14 +250,12 @@ impl Slider { // 5. 计算中心位置 (对齐 Python 逻辑) // target_w, target_h 来自输入数组的维度 let (th, tw) = target.dim(); - let (center_x, center_y) = - image_processor::calculate_center(max_loc, tw as usize, th as usize); + let (center_x, center_y) = image_processor::calculate_center(max_loc, tw, th); // 打印调试信息,方便与 Python 对比 - // println!("Edge Match: max_val: {}, max_loc: {:?}", max_val, max_loc); - println!("-Rust Target Width (tw): {}", tw); - println!("-Rust Best Max Loc X: {}", max_loc.0); - println!("-Rust Final Center X: {}", center_x); + debug!("-Rust Target Width (tw): {}", tw); + debug!("-Rust Best Max Loc X: {}", max_loc.0); + debug!("-Rust Final Center X: {}", center_x); SlideResult { target: [center_x, center_y], target_x: center_x, @@ -272,3 +264,44 @@ impl Slider { } } } + +#[cfg(test)] +mod tests { + use super::*; + + fn paint_block(img: &mut DynamicImage, x: u32, y: u32, w: u32, h: u32) { + let mut luma = img.to_luma8(); + for yy in y..y + h { + for xx in x..x + w { + luma.put_pixel(xx, yy, Luma([255u8])); + } + } + *img = DynamicImage::ImageLuma8(luma); + } + + #[test] + fn slide_match_finds_block_center() { + let slider = Slider::new(); + let mut target = DynamicImage::new_luma8(10, 10); + paint_block(&mut target, 3, 3, 4, 4); + let mut background = DynamicImage::new_luma8(30, 30); + paint_block(&mut background, 8, 8, 4, 4); + + let res = slider.slide_match(&target, &background, true).unwrap(); + // 模板内白块位于 (3,3),因此最佳匹配原点为 (8-3, 8-3)=(5,5),中心为 (5+5, 5+5) + assert_eq!(res.target_x, 10); + assert_eq!(res.target_y, 10); + assert!((res.confidence - 1.0).abs() < 1e-3); + } + + #[test] + fn slide_comparison_identical_returns_zero() { + let slider = Slider::new(); + let mut img = DynamicImage::new_luma8(16, 16); + paint_block(&mut img, 4, 4, 4, 4); + + let res = slider.slide_comparison(&img, &img).unwrap(); + assert_eq!(res.target, [0, 0]); + assert_eq!(res.confidence, 0.0); + } +} diff --git a/ddddocr-core/src/traits.rs b/ddddocr-core/src/traits.rs index c4e66fa..ab16993 100644 --- a/ddddocr-core/src/traits.rs +++ b/ddddocr-core/src/traits.rs @@ -7,8 +7,11 @@ use std::path::Path; /// 查询模型输入/输出信息的接口。 pub trait Info { + /// 获取输入张量信息列表。 fn input_info(&self) -> crate::error::Result>; + /// 获取输出张量信息列表。 fn output_info(&self) -> crate::error::Result>; + /// 获取模型完整输入/输出信息。 fn model_info(&self) -> crate::error::Result; } @@ -16,6 +19,7 @@ pub trait Info { pub trait InferenceEngine { /// 引擎产出的输出枚举(OCR 为 [`crate::OcrOutput`],检测为 [`crate::DetOutput`])。 type Output; + /// 对输入张量执行推理并返回引擎定义的输出。 fn inference( &self, input_array: ndarray::Array4, @@ -24,6 +28,7 @@ pub trait InferenceEngine { /// OCR 引擎接口:输出 [`crate::OcrOutput`],并提供模型元数据。 pub trait OcrEngine: InferenceEngine + Info { + /// 获取模型元数据。 fn metadata(&self) -> &ModelMetadata; } @@ -32,13 +37,17 @@ pub trait DetEngine: InferenceEngine {} /// 模型加载器:从本地路径或字节流构建引擎会话。 pub trait Loader { + /// 构建出的引擎会话类型。 type Session; + /// 构建过程中的错误类型。 type Error; + /// 从本地模型路径构建会话。 fn build_for_path>( &self, model_path: P, ) -> crate::error::Result; + /// 从模型字节流构建会话。 fn build_from_bytes( &self, model_bytes: &[u8], diff --git a/ddddocr-core/src/types.rs b/ddddocr-core/src/types.rs index 36ab7e8..7de1804 100644 --- a/ddddocr-core/src/types.rs +++ b/ddddocr-core/src/types.rs @@ -3,19 +3,25 @@ /// 张量元素的数据类型标记。 #[derive(Debug, Clone)] pub enum TensorType { + /// 32 位浮点类型。 F32, + /// 64 位整数类型。 I64, + /// 其他类型。 Other, } /// 模型某个轴的维度特征:静态数值或动态符号。 #[derive(Clone, PartialEq, Eq)] pub enum AxisDim { + /// 静态固定维度。 Static(usize), + /// 动态符号维度。 Dynamic(String), } impl AxisDim { + /// 是否为动态维度。 pub fn is_dynamic(&self) -> bool { matches!(self, AxisDim::Dynamic(_)) } @@ -34,15 +40,20 @@ impl std::fmt::Debug for AxisDim { /// 单个张量(输入或输出)的名称、形状与数据类型描述。 #[derive(Debug, Clone)] pub struct TensorInfo { + /// 张量名称。 pub name: String, + /// 各轴形状描述。 pub shape: Vec, + /// 元素数据类型。 pub tensor_type: TensorType, } /// 模型的完整输入/输出信息。 #[derive(Debug, Clone)] pub struct ModelInfo { + /// 输入张量列表。 pub inputs: Vec, + /// 输出张量列表。 pub outputs: Vec, /// 硬件执行提供者(`None` 表示使用引擎默认后端)。 pub providers: Option>, diff --git a/ddddocr-core/src/utils/image_convert.rs b/ddddocr-core/src/utils/image_convert.rs index 27eef47..ab18acb 100644 --- a/ddddocr-core/src/utils/image_convert.rs +++ b/ddddocr-core/src/utils/image_convert.rs @@ -7,21 +7,22 @@ use ndarray::{Array3, ArrayViewD}; /// 图像通道模式。 #[derive(Debug)] pub enum ColorMode { + /// RGB 三通道。 RGB, + /// RGBA 四通道。 RGBA, + /// 灰度单通道。 L, } /// 将 ndarray 数组转换为图像(自动识别 HWC 通道数)。 // 对应 Python 版 _numpy_to_pil_image -pub fn ndarray_to_hwc_image(array: ArrayViewD) -> Result { +pub fn ndarray_to_hwc_image(array: ArrayViewD) -> Result { let shape = array.shape(); let dim = shape.len(); // 1. 确保数据在内存中是连续的 (C order / Standard Layout) // 如果 arr 是经过切片或转置的,这一步会进行必要的内存拷贝 - // let standard = array.as_standard_layout(); - // let (raw_data, _offset) = standard.to_owned().into_raw_vec_and_offset(); let color_mode = match dim { // 对应 Python: len(array.shape) == 2 (灰度图 H, W) @@ -100,8 +101,10 @@ pub fn png_rgba_white_preprocess(img: &DynamicImage) -> DynamicImage { DynamicImage::ImageRgb8(background) } /// 将 DynamicImage 转换为 array 数组 -pub fn image_to_ndarray(image: &DynamicImage, mode: ColorMode) -> Result,ImagePreprocessError> { - // 1. 模式转换 (对应 utils.convert(target_mode)),此函数在时保留看后续优化是否需要替代image_to_ndarray +pub fn image_to_ndarray( + image: &DynamicImage, + mode: ColorMode, +) -> Result, ImagePreprocessError> { // Rust utils 库通过 to_rgb8, to_luma8 等方法实现转换 let (width, height) = image.dimensions(); @@ -116,7 +119,10 @@ pub fn image_to_ndarray(image: &DynamicImage, mode: ColorMode) -> Result, mode: ColorMode) -> Result { +pub fn ndarray_to_image( + array: ArrayViewD, + mode: ColorMode, +) -> Result { let shape = array.shape(); // 基础边界检查:至少要有 H 和 W 两个维度 @@ -124,12 +130,15 @@ pub fn ndarray_to_image(array: ArrayViewD, mode: ColorMode) -> Result, mode: ColorMode) -> Result { +fn from_ndarray( + array: ArrayViewD, + mode: ColorMode, +) -> Result { let shape = array.shape(); // 映射:ndarray 的 shape 默认是 [Height, Width, (Channels)] @@ -152,14 +161,12 @@ fn from_ndarray(array: ArrayViewD, mode: ColorMode) -> Result, mode: ColorMode) -> Result, b: &ArrayView3) -> Array3 { // 或者直接使用 zip_mut_with 处理以减少内存分配 let mut diff = Array3::zeros(a.dim()); azip!((res in &mut diff, &va in a, &vb in b) { - *res = (va as i16 - vb as i16).abs() as u8; + *res = va.abs_diff(vb); }); diff } @@ -190,3 +190,43 @@ pub fn resize_image( image.resize_exact(target_width, target_height, FilterType::Lanczos3) } +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn abs_diff_absolutes() { + let a = Array3::from_shape_vec((1, 1, 3), vec![10, 200, 5]).unwrap(); + let b = Array3::from_shape_vec((1, 1, 3), vec![200, 100, 5]).unwrap(); + let d = abs_diff(&a.view(), &b.view()); + assert_eq!(d[[0, 0, 0]], 190); + assert_eq!(d[[0, 0, 1]], 100); + assert_eq!(d[[0, 0, 2]], 0); + } + + #[test] + fn rgb_to_gray_white_is_255() { + let rgb = Array3::from_shape_vec((1, 1, 3), vec![255, 255, 255]).unwrap(); + assert_eq!(rgb_to_gray(rgb.view())[[0, 0]], 255); + } + + #[test] + fn min_max_loc_finds_max() { + let buf = ImageBuffer::, Vec>::from_fn(3, 2, |x, y| { + Luma([if x == 2 && y == 1 { 0.9 } else { 0.1 }]) + }); + let (val, loc) = min_max_loc(&buf); + assert_eq!(val, 0.9); + assert_eq!(loc, (2, 1)); + } + + #[test] + fn calculate_center_midpoint() { + assert_eq!(calculate_center((10, 20), 6, 4), (13, 22)); + } + + #[test] + fn rgb_to_opencv_hsv_red() { + assert_eq!(rgb_to_opencv_hsv(255, 0, 0), (0, 255, 255)); + } +} diff --git a/ddddocr-core/src/utils/tensor_transform.rs b/ddddocr-core/src/utils/tensor_transform.rs index 4b45bb0..981edc6 100644 --- a/ddddocr-core/src/utils/tensor_transform.rs +++ b/ddddocr-core/src/utils/tensor_transform.rs @@ -5,7 +5,10 @@ use crate::error::{Result, TensorError}; use ndarray::s; /// 将异构形状的模型输出规整为标准 `[Steps, Classes]` Logits 矩阵。 -pub fn normalize_ocr_logits(array: ndarray::ArrayViewD, shape: &[usize]) -> Result { +pub fn normalize_ocr_logits( + array: ndarray::ArrayViewD, + shape: &[usize], +) -> Result { let (steps, classes, data_dyn_view) = match shape.len() { 3 => { if shape[1] == 1 { @@ -27,12 +30,10 @@ pub fn normalize_ocr_logits(array: ndarray::ArrayViewD, shape: &[usize]) -> // 我们把它虚构成一个 [1, Classes] 的 2D 矩阵来复用后面的 argmax 逻辑 1 => (1, shape[0], array), _ => { - return Err( - TensorError::DimensionMismatch { - expected: "1D, 2D, or 3D OCR Logits".to_string(), - actual: shape.to_vec(), - }, - ); + return Err(TensorError::DimensionMismatch { + expected: "1D, 2D, or 3D OCR Logits".to_string(), + actual: shape.to_vec(), + }); } }; @@ -52,3 +53,32 @@ pub fn normalize_ocr_logits(array: ndarray::ArrayViewD, shape: &[usize]) -> Ok(OcrOutput::Logits(matrix_cow)) } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn normalize_2d_logits() { + let array = ndarray::Array2::::zeros((2, 3)); + match normalize_ocr_logits(array.view().into_dyn(), &[2, 3]).unwrap() { + crate::OcrOutput::Logits(m) => assert_eq!(m.dim(), (2, 3)), + _ => panic!("应为 Logits 输出"), + } + } + + #[test] + fn normalize_3d_batch_first() { + let array = ndarray::Array3::::zeros((1, 2, 3)); + match normalize_ocr_logits(array.view().into_dyn(), &[1, 2, 3]).unwrap() { + crate::OcrOutput::Logits(m) => assert_eq!(m.dim(), (2, 3)), + _ => panic!("应为 Logits 输出"), + } + } + + #[test] + fn unsupported_dim_errors() { + let array = ndarray::Array4::::zeros((1, 1, 1, 1)); + assert!(normalize_ocr_logits(array.view().into_dyn(), &[1, 1, 1, 1]).is_err()); + } +} diff --git a/ddddocr-core/tests/api_surface.rs b/ddddocr-core/tests/api_surface.rs new file mode 100644 index 0000000..1416a6f --- /dev/null +++ b/ddddocr-core/tests/api_surface.rs @@ -0,0 +1,72 @@ +//! 外部视角 API 测试:验证颜色过滤与字符集限制扩展点对外可用。 + +use std::borrow::Cow; + +use ddddocr_core::color_any_of; +use ddddocr_core::{ + CharRestrict, ColorFilter, ColorPreset, HsvRange, IdRestrict, OcrBuilder, TokenFilter, any_of, +}; + +#[test] +fn color_filter_setter_accepts_owned_preset() { + let builder = OcrBuilder::new().color_filter(ColorPreset::Red); + let _ = builder; +} + +#[test] +fn color_any_of_macro_expands_and_collects() { + let ranges = color_any_of!(ColorPreset::Red, ColorPreset::Blue) + .collect_to_vec() + .expect("收集颜色区间失败") + .expect("应存在有效颜色区间"); + // Red 两段 + Blue 一段 + assert_eq!(ranges.len(), 3); +} + +#[test] +fn charset_restrict_setter_accepts_owned_restrict() { + let builder = OcrBuilder::new().charset_restrict(CharRestrict::Digit); + let _ = builder; +} + +#[test] +fn any_of_macro_expands_and_filters() { + let tokens: Vec> = vec![ + Cow::Borrowed(""), + Cow::Borrowed("1"), + Cow::Borrowed("a"), + Cow::Borrowed("A"), + ]; + + let indices = any_of!(CharRestrict::Digit, CharRestrict::Lowercase) + .apply_to_charset(&tokens) + .expect("字符集应存在交集"); + // 0 号 blank 放行 + 数字 '1' + 小写 'a' + assert_eq!(indices, vec![0, 1, 2]); +} + +#[test] +fn id_restrict_filters_by_index() { + let tokens: Vec> = vec![ + Cow::Borrowed(""), + Cow::Borrowed("1"), + Cow::Borrowed("2"), + Cow::Borrowed("3"), + ]; + + let indices = IdRestrict::TopN(2) + .apply_to_charset(&tokens) + .expect("字符集应存在交集"); + assert_eq!(indices, vec![0, 1]); +} + +#[test] +fn hsv_range_validate() { + let range = HsvRange::new((0, 50, 50), (10, 255, 255)); + assert!(range.validate().is_ok()); + assert!( + HsvRange::new((200, 0, 0), (255, 255, 255)) + .validate() + .is_err() + ); +}