feat(core): 扩展 API、完善日志与代码文档规范

- 公开颜色过滤与字符集限制扩展 API,修复宏路径
- 库内打印替换为 tracing 日志,清理遗留废弃代码
- 补充核心逻辑单元测试与 crate 元数据
- 开启 missing_docs 并统一 rustfmt/clippy 格式
This commit is contained in:
2026-08-06 19:58:54 +08:00
parent 1362243f4e
commit fe61895926
20 changed files with 696 additions and 202 deletions

View File

@@ -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}
thiserror = { workspace = true }
tracing = { workspace = true }

View File

@@ -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<Vec<TensorInfo>> {
Ok(vec![])
}
fn output_info(&self) -> Result<Vec<TensorInfo>> {
Ok(vec![])
}
fn model_info(&self) -> Result<ModelInfo> {
Ok(ModelInfo {
inputs: vec![],
outputs: vec![],
providers: None,
})
}
}
impl InferenceEngine for DemoEngine {
type Output = OcrOutput;
fn inference(&self, input: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
// 用全零 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}");
}

View File

@@ -1,7 +1,6 @@
//! 检测器构建器。
use crate::det::executor::Detector;
// use ddddocr_tract::det::session::DetSession;
use crate::traits::DetEngine;
/// 检测器构建器,通过 [`crate::Detector::builder`] 创建。

View File

@@ -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<Vec<DetectionResult>, 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::<f32>::zeros(boxes.raw_dim());
for i in 0..boxes.nrows() {

View File

@@ -3,18 +3,20 @@
use thiserror::Error;
/// 全局统一的 `Result` 别名,默认错误类型为 [`DdddError`]。
// pub type Result<T> = std::result::Result<T, DdddError>;
pub type Result<T, E = DdddError> = std::result::Result<T, E>;
/// 顶层错误类型,聚合本库各阶段错误。
#[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<usize>,
},
/// 图像缓冲区长度与分辨率/通道数不匹配。
#[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<usize>,
},
/// 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(_)));
}
}

View File

@@ -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<i64>),
/// Logits 矩阵输出 `[Steps, Classes]`。
Logits(ndarray::Array2<f32>),
}
/// 目标检测模型的统一输出枚举,由推理引擎产出,供 [`Detector`] 后处理。
pub enum DetOutput {
/// 原始检测输出张量。
Detection(ndarray::Array3<f32>),
}

View File

@@ -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<T>(mut self, filter: T) -> Self
where
T: ColorFilter + Send + Sync + 'static,
@@ -45,6 +49,7 @@ impl OcrBuilder {
self
}
/// 设置字符集限制。
pub fn charset_restrict<T>(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, // 原地解构出来

View File

@@ -9,22 +9,22 @@ use std::collections::HashMap;
/// 字符集token 列表与索引的双向映射。
#[derive(Debug, Clone)]
pub struct Charset {
/// 字符集 token 列表。
// 使用 Cow 统一静态切片和动态读取的 Vec<String>,内部实现真正的零拷贝
pub tokens: Vec<Cow<'static, str>>,
/// 字符到索引的反查表。
// 反向查找表,保证字符转索引为 O(1)
pub char_to_idx: HashMap<Cow<'static, str>, usize>,
// 当前处于激活状态的有效索引缓存 (用于 CTC 解码前的过滤加速)
// pub valid_indices: HashSet<usize>,
}
impl Charset {
// 内部底层统一收拢构造
/// 从 token 列表构建字符集。
pub fn new(tokens: Vec<Cow<'static, str>>) -> 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<Cow<'static, str>> {
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);
}
}

View File

@@ -48,8 +48,7 @@ pub fn apply_to_image(
// 3. 将扁平字节数组重新打包回 DynamicImage 容器
let filtered_buffer = ImageBuffer::<Rgb<u8>, Vec<u8>>::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<HsvRange>),
}
@@ -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<HsvRange>);
/// 预估有效区间数量。
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::<ColorPreset>().unwrap(), ColorPreset::Red);
assert!("pink".parse::<ColorPreset>().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]);
}
}
}

View File

@@ -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<Vec<f32>>,
@@ -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<OcrResult> {
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<OcrResult, TensorError> {
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<f32>,且保证是 [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<OcrResult> {
// // 1. 拿到底层的动态维度只读视图
// let view = raw_tensor.to_array_view::<i64>()?;
//
// // 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<OcrResult> {
// let shape = raw_tensor.shape();
// println!("模型输出shape数据: {:?}", shape);
// let view = raw_tensor.to_array_view::<f32>()?;
//
// // 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<f32> = 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<i64> = 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. 字符映射

View File

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

View File

@@ -1,15 +1,19 @@
//! 字符集限制:按字符属性或索引过滤识别范围。
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 {
@@ -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<String>),
}
@@ -90,8 +98,11 @@ impl TokenFilter for CharRestrict {
/// 按索引限制:前 N 个、索引范围或索引列表。
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IdRestrict {
/// 前 N 个索引。
TopN(usize),
/// 指定索引范围。
IdRange(std::ops::Range<usize>),
/// 指定索引列表。
IdList(Vec<usize>),
}
@@ -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<Cow<'static, str>> {
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]);
}
}

View File

@@ -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<u8>,
@@ -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);
}
}

View File

@@ -7,8 +7,11 @@ use std::path::Path;
/// 查询模型输入/输出信息的接口。
pub trait Info {
/// 获取输入张量信息列表。
fn input_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
/// 获取输出张量信息列表。
fn output_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
/// 获取模型完整输入/输出信息。
fn model_info(&self) -> crate::error::Result<ModelInfo>;
}
@@ -16,6 +19,7 @@ pub trait Info {
pub trait InferenceEngine {
/// 引擎产出的输出枚举OCR 为 [`crate::OcrOutput`],检测为 [`crate::DetOutput`])。
type Output;
/// 对输入张量执行推理并返回引擎定义的输出。
fn inference(
&self,
input_array: ndarray::Array4<f32>,
@@ -24,6 +28,7 @@ pub trait InferenceEngine {
/// OCR 引擎接口:输出 [`crate::OcrOutput`],并提供模型元数据。
pub trait OcrEngine: InferenceEngine<Output = OcrOutput> + Info {
/// 获取模型元数据。
fn metadata(&self) -> &ModelMetadata;
}
@@ -32,13 +37,17 @@ pub trait DetEngine: InferenceEngine<Output = DetOutput> {}
/// 模型加载器:从本地路径或字节流构建引擎会话。
pub trait Loader {
/// 构建出的引擎会话类型。
type Session;
/// 构建过程中的错误类型。
type Error;
/// 从本地模型路径构建会话。
fn build_for_path<P: AsRef<Path>>(
&self,
model_path: P,
) -> crate::error::Result<Self::Session, Self::Error>;
/// 从模型字节流构建会话。
fn build_from_bytes(
&self,
model_bytes: &[u8],

View File

@@ -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<AxisDim>,
/// 元素数据类型。
pub tensor_type: TensorType,
}
/// 模型的完整输入/输出信息。
#[derive(Debug, Clone)]
pub struct ModelInfo {
/// 输入张量列表。
pub inputs: Vec<TensorInfo>,
/// 输出张量列表。
pub outputs: Vec<TensorInfo>,
/// 硬件执行提供者(`None` 表示使用引擎默认后端)。
pub providers: Option<Vec<String>>,

View File

@@ -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<u8>) -> Result<DynamicImage,ImagePreprocessError> {
pub fn ndarray_to_hwc_image(array: ArrayViewD<u8>) -> Result<DynamicImage, ImagePreprocessError> {
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<Array3<u8>,ImagePreprocessError> {
// 1. 模式转换 (对应 utils.convert(target_mode)),此函数在时保留看后续优化是否需要替代image_to_ndarray
pub fn image_to_ndarray(
image: &DynamicImage,
mode: ColorMode,
) -> Result<Array3<u8>, 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<Array3<
Ok(array)
}
/// 将 array 数组转换为 DynamicImage
pub fn ndarray_to_image(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,ImagePreprocessError> {
pub fn ndarray_to_image(
array: ArrayViewD<u8>,
mode: ColorMode,
) -> Result<DynamicImage, ImagePreprocessError> {
let shape = array.shape();
// 基础边界检查:至少要有 H 和 W 两个维度
@@ -124,12 +130,15 @@ pub fn ndarray_to_image(array: ArrayViewD<u8>, mode: ColorMode) -> Result<Dynami
return Err(ImagePreprocessError::InvalidDimensions {
expected: "至少为 2D array [H, W]".to_string(),
actual: shape.to_vec(),
})?;
});
}
from_ndarray(array, mode)
}
fn from_ndarray(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,ImagePreprocessError> {
fn from_ndarray(
array: ArrayViewD<u8>,
mode: ColorMode,
) -> Result<DynamicImage, ImagePreprocessError> {
let shape = array.shape();
// 映射ndarray 的 shape 默认是 [Height, Width, (Channels)]
@@ -152,14 +161,12 @@ fn from_ndarray(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,I
let expected_len = (width * height) as usize * channels;
// 构造通用错误闭包,避免 match 分支中重复编写冗长的错误对象
let make_err = || {
ImagePreprocessError::BufferLengthMismatch {
expected: expected_len,
actual: raw_len,
width,
height,
channels,
}
let make_err = || ImagePreprocessError::BufferLengthMismatch {
expected: expected_len,
actual: raw_len,
width,
height,
channels,
};
// 2. 重新解释内存并构建 ImageBuffer
@@ -175,3 +182,36 @@ fn from_ndarray(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,I
.ok_or_else(make_err),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn image_to_ndarray_rgb_dims() {
let img = DynamicImage::new_rgb8(4, 3);
let arr = image_to_ndarray(&img, ColorMode::RGB).unwrap();
assert_eq!(arr.dim(), (3, 4, 3));
}
#[test]
fn ndarray_roundtrip() {
let img = DynamicImage::new_rgb8(2, 2);
let arr = image_to_ndarray(&img, ColorMode::RGB).unwrap();
let back = ndarray_to_image(arr.view().into_dyn(), ColorMode::RGB).unwrap();
assert_eq!(back.to_rgb8().dimensions(), (2, 2));
}
#[test]
fn png_white_preprocess_fills_transparent() {
let img = DynamicImage::ImageRgba8(image::RgbaImage::from_pixel(
2,
2,
image::Rgba([255, 0, 0, 0]),
));
let out = png_rgba_white_preprocess(&img);
for p in out.to_rgb8().pixels() {
assert_eq!(p.0, [255, 255, 255]);
}
}
}

View File

@@ -2,7 +2,7 @@
use crate::error::{DdddError, Result};
use crate::utils::image_convert::ndarray_to_hwc_image;
use base64::{engine::general_purpose, Engine as _};
use base64::{Engine as _, engine::general_purpose};
use image::DynamicImage;
use ndarray::ArrayViewD;
use std::fmt;

View File

@@ -1,8 +1,8 @@
//! 图像处理算法OpenCV 风格的常用函数封装。
use image::{imageops::FilterType, DynamicImage, GrayImage, ImageBuffer, Luma};
use image::{DynamicImage, GrayImage, ImageBuffer, Luma, imageops::FilterType};
use ndarray::{azip, Array2, Array3, ArrayView2, ArrayView3};
use ndarray::{Array2, Array3, ArrayView2, ArrayView3, azip};
use std::cmp::{max, min};
// 模拟openCV
@@ -12,7 +12,7 @@ pub fn abs_diff(a: &ArrayView3<u8>, b: &ArrayView3<u8>) -> Array3<u8> {
// 或者直接使用 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::<Luma<f32>, Vec<f32>>::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));
}
}

View File

@@ -5,7 +5,10 @@ use crate::error::{Result, TensorError};
use ndarray::s;
/// 将异构形状的模型输出规整为标准 `[Steps, Classes]` Logits 矩阵。
pub fn normalize_ocr_logits(array: ndarray::ArrayViewD<f32>, shape: &[usize]) -> Result<OcrOutput,TensorError> {
pub fn normalize_ocr_logits(
array: ndarray::ArrayViewD<f32>,
shape: &[usize],
) -> Result<OcrOutput, TensorError> {
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<f32>, 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<f32>, shape: &[usize]) ->
Ok(OcrOutput::Logits(matrix_cow))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn normalize_2d_logits() {
let array = ndarray::Array2::<f32>::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::<f32>::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::<f32>::zeros((1, 1, 1, 1));
assert!(normalize_ocr_logits(array.view().into_dyn(), &[1, 1, 1, 1]).is_err());
}
}

View File

@@ -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<Cow<'static, str>> = 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<Cow<'static, str>> = 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()
);
}