feat(core): 扩展 API、完善日志与代码文档规范
- 公开颜色过滤与字符集限制扩展 API,修复宏路径 - 库内打印替换为 tracing 日志,清理遗留废弃代码 - 补充核心逻辑单元测试与 crate 元数据 - 开启 missing_docs 并统一 rustfmt/clippy 格式
This commit is contained in:
@@ -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 }
|
||||
|
||||
63
ddddocr-core/examples/quick_start.rs
Normal file
63
ddddocr-core/examples/quick_start.rs
Normal 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}");
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
//! 检测器构建器。
|
||||
|
||||
use crate::det::executor::Detector;
|
||||
// use ddddocr_tract::det::session::DetSession;
|
||||
use crate::traits::DetEngine;
|
||||
|
||||
/// 检测器构建器,通过 [`crate::Detector::builder`] 创建。
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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(_)));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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>),
|
||||
}
|
||||
|
||||
@@ -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, // 原地解构出来
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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. 字符映射
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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>>,
|
||||
|
||||
@@ -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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
72
ddddocr-core/tests/api_surface.rs
Normal file
72
ddddocr-core/tests/api_surface.rs
Normal 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()
|
||||
);
|
||||
}
|
||||
Reference in New Issue
Block a user