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 }
|
version = { workspace = true }
|
||||||
edition = { workspace = true }
|
edition = { workspace = true }
|
||||||
license = { 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]
|
[dependencies]
|
||||||
ndarray = { workspace = true } # 继承自工作空间
|
ndarray = { workspace = true }
|
||||||
|
|
||||||
base64 = { workspace = true }
|
base64 = { workspace = true }
|
||||||
|
|
||||||
image = { workspace = true }
|
image = { workspace = true }
|
||||||
imageproc = { workspace = true }
|
imageproc = { workspace = true }
|
||||||
|
thiserror = { workspace = true }
|
||||||
thiserror = { workspace = true } # 刚好可以开始接入你需要的标准库错误处理
|
tracing = { 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 crate::det::executor::Detector;
|
||||||
// use ddddocr_tract::det::session::DetSession;
|
|
||||||
use crate::traits::DetEngine;
|
use crate::traits::DetEngine;
|
||||||
|
|
||||||
/// 检测器构建器,通过 [`crate::Detector::builder`] 创建。
|
/// 检测器构建器,通过 [`crate::Detector::builder`] 创建。
|
||||||
|
|||||||
@@ -11,11 +11,17 @@ use crate::{DetBuilder, DetOutput};
|
|||||||
/// 目标检测结果:原图像素坐标系下的边界框、置信度与类别 ID。
|
/// 目标检测结果:原图像素坐标系下的边界框、置信度与类别 ID。
|
||||||
#[derive(Debug, Clone, Copy)]
|
#[derive(Debug, Clone, Copy)]
|
||||||
pub struct DetectionResult {
|
pub struct DetectionResult {
|
||||||
|
/// 左上角 x 坐标。
|
||||||
pub x1: i32,
|
pub x1: i32,
|
||||||
|
/// 左上角 y 坐标。
|
||||||
pub y1: i32,
|
pub y1: i32,
|
||||||
|
/// 右下角 x 坐标。
|
||||||
pub x2: i32,
|
pub x2: i32,
|
||||||
|
/// 右下角 y 坐标。
|
||||||
pub y2: i32,
|
pub y2: i32,
|
||||||
|
/// 置信度。
|
||||||
pub score: f32,
|
pub score: f32,
|
||||||
|
/// 类别 ID。
|
||||||
pub class_id: u32,
|
pub class_id: u32,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -36,11 +42,13 @@ pub struct Detector<'a> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl<'a> Detector<'a> {
|
impl<'a> Detector<'a> {
|
||||||
|
/// 绑定检测引擎会话创建检测器。
|
||||||
pub fn new(runtime: &'a dyn DetEngine) -> Self {
|
pub fn new(runtime: &'a dyn DetEngine) -> Self {
|
||||||
Detector { runtime }
|
Detector { runtime }
|
||||||
}
|
}
|
||||||
|
/// 创建检测器构建器。
|
||||||
pub fn builder() -> DetBuilder {
|
pub fn builder() -> DetBuilder {
|
||||||
DetBuilder::default()
|
DetBuilder
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
impl<'a> Detector<'a> {
|
impl<'a> Detector<'a> {
|
||||||
@@ -247,15 +255,12 @@ impl<'a> Detector<'a> {
|
|||||||
dynamic_img: &DynamicImage,
|
dynamic_img: &DynamicImage,
|
||||||
) -> Result<Vec<DetectionResult>, TensorError> {
|
) -> Result<Vec<DetectionResult>, TensorError> {
|
||||||
// 使用 utils crate 解码
|
// 使用 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 (orig_w, orig_h) = dynamic_img.dimensions();
|
||||||
|
|
||||||
let (input_tensor, ratio) = self.preproc(dynamic_img, (416, 416));
|
let (input_tensor, ratio) = self.preproc(dynamic_img, (416, 416));
|
||||||
|
|
||||||
// tract 推理
|
// tract 推理
|
||||||
// let outputs = self.session.session.run(tvec!(input_tensor.into()))?;
|
|
||||||
let outputs = self.runtime.inference(input_tensor)?;
|
let outputs = self.runtime.inference(input_tensor)?;
|
||||||
// let output_array = outputs[0]
|
|
||||||
// 2. 无缝、安全地解包出标准 3维 矩阵
|
// 2. 无缝、安全地解包出标准 3维 矩阵
|
||||||
let DetOutput::Detection(output_array) = outputs;
|
let DetOutput::Detection(output_array) = outputs;
|
||||||
|
|
||||||
@@ -272,9 +277,7 @@ impl<'a> Detector<'a> {
|
|||||||
expected: format!("可广播至 cls_conf 形状 {:?}", cls_conf.shape()),
|
expected: format!("可广播至 cls_conf 形状 {:?}", cls_conf.shape()),
|
||||||
actual: obj_conf.shape().to_vec(),
|
actual: obj_conf.shape().to_vec(),
|
||||||
})?;
|
})?;
|
||||||
// .context("ndarray broadcasting failed for scores calculation")?;
|
|
||||||
let scores = &obj_broadcast * &cls_conf;
|
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());
|
let mut boxes_xyxy = Array2::<f32>::zeros(boxes.raw_dim());
|
||||||
for i in 0..boxes.nrows() {
|
for i in 0..boxes.nrows() {
|
||||||
|
|||||||
@@ -3,18 +3,20 @@
|
|||||||
use thiserror::Error;
|
use thiserror::Error;
|
||||||
|
|
||||||
/// 全局统一的 `Result` 别名,默认错误类型为 [`DdddError`]。
|
/// 全局统一的 `Result` 别名,默认错误类型为 [`DdddError`]。
|
||||||
// pub type Result<T> = std::result::Result<T, DdddError>;
|
|
||||||
pub type Result<T, E = DdddError> = std::result::Result<T, E>;
|
pub type Result<T, E = DdddError> = std::result::Result<T, E>;
|
||||||
|
|
||||||
/// 顶层错误类型,聚合本库各阶段错误。
|
/// 顶层错误类型,聚合本库各阶段错误。
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
pub enum DdddError {
|
pub enum DdddError {
|
||||||
|
/// 图像预处理阶段异常。
|
||||||
#[error("图像预处理失败: {0}")]
|
#[error("图像预处理失败: {0}")]
|
||||||
Preprocess(#[from] ImagePreprocessError),
|
Preprocess(#[from] ImagePreprocessError),
|
||||||
|
|
||||||
|
/// 推理与张量操作阶段异常。
|
||||||
#[error("推理与模型输入/输出张量异常: {0}")]
|
#[error("推理与模型输入/输出张量异常: {0}")]
|
||||||
Inference(#[from] TensorError),
|
Inference(#[from] TensorError),
|
||||||
|
|
||||||
|
/// 后处理解码阶段异常。
|
||||||
#[error("后处理解码错误: {0}")]
|
#[error("后处理解码错误: {0}")]
|
||||||
Decode(#[from] DecodeError),
|
Decode(#[from] DecodeError),
|
||||||
|
|
||||||
@@ -30,46 +32,66 @@ pub enum DdddError {
|
|||||||
/// 图像预处理阶段错误类型。
|
/// 图像预处理阶段错误类型。
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
pub enum ImagePreprocessError {
|
pub enum ImagePreprocessError {
|
||||||
|
/// ndarray 基础操作失败。
|
||||||
#[error("图片转矩阵(ndarray)基础操作失败: {0}")]
|
#[error("图片转矩阵(ndarray)基础操作失败: {0}")]
|
||||||
Ndarray(#[from] ndarray::ShapeError),
|
Ndarray(#[from] ndarray::ShapeError),
|
||||||
|
|
||||||
|
/// 图像矩阵维度不合规。
|
||||||
#[error("图像矩阵维度不合规!预期: {expected},实际图像形状: {actual:?}")]
|
#[error("图像矩阵维度不合规!预期: {expected},实际图像形状: {actual:?}")]
|
||||||
InvalidDimensions {
|
InvalidDimensions {
|
||||||
|
/// 期望的维度描述。
|
||||||
expected: String,
|
expected: String,
|
||||||
|
/// 实际的图像形状。
|
||||||
actual: Vec<usize>,
|
actual: Vec<usize>,
|
||||||
},
|
},
|
||||||
|
|
||||||
|
/// 图像缓冲区长度与分辨率/通道数不匹配。
|
||||||
#[error(
|
#[error(
|
||||||
"图像缓冲区长度不匹配!预期大小: {expected},实际大小: {actual} (分辨率: {width}x{height}, 通道数: {channels})"
|
"图像缓冲区长度不匹配!预期大小: {expected},实际大小: {actual} (分辨率: {width}x{height}, 通道数: {channels})"
|
||||||
)]
|
)]
|
||||||
BufferLengthMismatch {
|
BufferLengthMismatch {
|
||||||
|
/// 期望的缓冲区长度。
|
||||||
expected: usize,
|
expected: usize,
|
||||||
|
/// 实际的缓冲区长度。
|
||||||
actual: usize,
|
actual: usize,
|
||||||
|
/// 图像宽度。
|
||||||
width: u32,
|
width: u32,
|
||||||
|
/// 图像高度。
|
||||||
height: u32,
|
height: u32,
|
||||||
|
/// 图像通道数。
|
||||||
channels: usize,
|
channels: usize,
|
||||||
},
|
},
|
||||||
|
|
||||||
|
/// 不支持的图像通道数。
|
||||||
#[error("不支持的图像通道数: {0} (仅支持单通道灰度L、3通道RGB、4通道RGBA)")]
|
#[error("不支持的图像通道数: {0} (仅支持单通道灰度L、3通道RGB、4通道RGBA)")]
|
||||||
UnsupportedChannels(usize),
|
UnsupportedChannels(usize),
|
||||||
|
|
||||||
|
/// HSV 颜色区间参数非法。
|
||||||
#[error("HSV 颜色区间参数非法: {0}")]
|
#[error("HSV 颜色区间参数非法: {0}")]
|
||||||
InvalidHsvRange(String),
|
InvalidHsvRange(String),
|
||||||
|
|
||||||
|
/// 未知的颜色预设名称。
|
||||||
#[error("不支持的颜色预设名称: {0}")]
|
#[error("不支持的颜色预设名称: {0}")]
|
||||||
UnknownColorPreset(String),
|
UnknownColorPreset(String),
|
||||||
|
|
||||||
|
/// 颜色过滤器配置无效。
|
||||||
#[error("颜色过滤器配置无效或初始化失败: {0}")]
|
#[error("颜色过滤器配置无效或初始化失败: {0}")]
|
||||||
FilterConfigInvalid(String),
|
FilterConfigInvalid(String),
|
||||||
|
|
||||||
|
/// 图像维度不匹配。
|
||||||
#[error("图像维度不匹配!{0}")]
|
#[error("图像维度不匹配!{0}")]
|
||||||
MismatchDimensions(String),
|
MismatchDimensions(String),
|
||||||
|
|
||||||
|
/// 滑块模板尺寸大于背景图。
|
||||||
#[error("滑块模板尺寸 [{target_w}x{target_h}] 大于背景图 [{bg_w}x{bg_h}]")]
|
#[error("滑块模板尺寸 [{target_w}x{target_h}] 大于背景图 [{bg_w}x{bg_h}]")]
|
||||||
TargetExceedsBackground {
|
TargetExceedsBackground {
|
||||||
|
/// 滑块模板宽度。
|
||||||
target_w: usize,
|
target_w: usize,
|
||||||
|
/// 滑块模板高度。
|
||||||
target_h: usize,
|
target_h: usize,
|
||||||
|
/// 背景图宽度。
|
||||||
bg_w: usize,
|
bg_w: usize,
|
||||||
|
/// 背景图高度。
|
||||||
bg_h: usize,
|
bg_h: usize,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -77,21 +99,28 @@ pub enum ImagePreprocessError {
|
|||||||
/// 推理与张量操作阶段错误类型。
|
/// 推理与张量操作阶段错误类型。
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
pub enum TensorError {
|
pub enum TensorError {
|
||||||
|
/// 推理引擎内部异常。
|
||||||
#[error("推理引擎内部发生异常: {0}")]
|
#[error("推理引擎内部发生异常: {0}")]
|
||||||
Engine(String),
|
Engine(String),
|
||||||
|
|
||||||
|
/// 模型张量维度不匹配。
|
||||||
#[error("模型张量维度不匹配!预期: {expected},实际 Tensor 形状: {actual:?}")]
|
#[error("模型张量维度不匹配!预期: {expected},实际 Tensor 形状: {actual:?}")]
|
||||||
DimensionMismatch {
|
DimensionMismatch {
|
||||||
|
/// 期望的维度描述。
|
||||||
expected: String,
|
expected: String,
|
||||||
|
/// 实际的 Tensor 形状。
|
||||||
actual: Vec<usize>,
|
actual: Vec<usize>,
|
||||||
},
|
},
|
||||||
|
|
||||||
|
/// OCR Logits 矩阵变形失败。
|
||||||
#[error("OCR Logits 矩阵变形失败: {0}")]
|
#[error("OCR Logits 矩阵变形失败: {0}")]
|
||||||
LogitsDimensionMismatch(#[from] ndarray::ShapeError),
|
LogitsDimensionMismatch(#[from] ndarray::ShapeError),
|
||||||
|
|
||||||
|
/// 张量内存不连续。
|
||||||
#[error("内存不连续,无法执行零拷贝操作")]
|
#[error("内存不连续,无法执行零拷贝操作")]
|
||||||
NonContiguousMemory,
|
NonContiguousMemory,
|
||||||
|
|
||||||
|
/// 未知的模型输出格式。
|
||||||
#[error("未知的模型输出格式")]
|
#[error("未知的模型输出格式")]
|
||||||
UnknownOutputFormat,
|
UnknownOutputFormat,
|
||||||
}
|
}
|
||||||
@@ -99,6 +128,7 @@ pub enum TensorError {
|
|||||||
/// 算法解码阶段错误类型。
|
/// 算法解码阶段错误类型。
|
||||||
#[derive(Error, Debug)]
|
#[derive(Error, Debug)]
|
||||||
pub enum DecodeError {
|
pub enum DecodeError {
|
||||||
|
/// CTC 解码异常。
|
||||||
#[error("CTC 解码异常: {0}")]
|
#[error("CTC 解码异常: {0}")]
|
||||||
Ctc(String),
|
Ctc(String),
|
||||||
}
|
}
|
||||||
@@ -112,6 +142,7 @@ impl DdddError {
|
|||||||
DdddError::Other(error.into())
|
DdddError::Other(error.into())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 是否为图片维度不合规错误。
|
||||||
pub fn is_invalid_dimensions(&self) -> bool {
|
pub fn is_invalid_dimensions(&self) -> bool {
|
||||||
matches!(
|
matches!(
|
||||||
self,
|
self,
|
||||||
@@ -119,6 +150,7 @@ impl DdddError {
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 是否因通道数不合规而失败。
|
||||||
pub fn is_unsupported_channels(&self) -> bool {
|
pub fn is_unsupported_channels(&self) -> bool {
|
||||||
matches!(
|
matches!(
|
||||||
self,
|
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`。
|
//! 完整可运行示例见 `ddddocr-core/examples/quick_start.rs`。
|
||||||
|
|
||||||
|
#![warn(missing_docs)]
|
||||||
|
|
||||||
mod det;
|
mod det;
|
||||||
/// 分层错误类型。
|
/// 分层错误类型。
|
||||||
pub mod error;
|
pub mod error;
|
||||||
@@ -28,11 +30,14 @@ pub use crate::slide::{SlideResult, Slider};
|
|||||||
|
|
||||||
/// OCR 模型的统一输出枚举,由推理引擎产出,供 [`Ocr`] 后处理。
|
/// OCR 模型的统一输出枚举,由推理引擎产出,供 [`Ocr`] 后处理。
|
||||||
pub enum OcrOutput {
|
pub enum OcrOutput {
|
||||||
|
/// 索引序列输出(CTC 解码输入)。
|
||||||
Indices(ndarray::Array1<i64>),
|
Indices(ndarray::Array1<i64>),
|
||||||
|
/// Logits 矩阵输出 `[Steps, Classes]`。
|
||||||
Logits(ndarray::Array2<f32>),
|
Logits(ndarray::Array2<f32>),
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 目标检测模型的统一输出枚举,由推理引擎产出,供 [`Detector`] 后处理。
|
/// 目标检测模型的统一输出枚举,由推理引擎产出,供 [`Detector`] 后处理。
|
||||||
pub enum DetOutput {
|
pub enum DetOutput {
|
||||||
|
/// 原始检测输出张量。
|
||||||
Detection(ndarray::Array3<f32>),
|
Detection(ndarray::Array3<f32>),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ pub struct OcrBuilder {
|
|||||||
|
|
||||||
impl OcrBuilder {
|
impl OcrBuilder {
|
||||||
// 初始化任务,设置默认参数
|
// 初始化任务,设置默认参数
|
||||||
|
/// 创建默认配置的构建器。
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self {
|
Self {
|
||||||
png_fix: false, // 默认值
|
png_fix: false, // 默认值
|
||||||
@@ -28,15 +29,18 @@ impl OcrBuilder {
|
|||||||
charset_restrict: None,
|
charset_restrict: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
/// 设置是否修复 PNG 透明背景问题。
|
||||||
pub fn png_fix(mut self, value: bool) -> Self {
|
pub fn png_fix(mut self, value: bool) -> Self {
|
||||||
self.png_fix = value;
|
self.png_fix = value;
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
/// 设置是否返回概率信息。
|
||||||
pub fn probability(mut self, value: bool) -> Self {
|
pub fn probability(mut self, value: bool) -> Self {
|
||||||
self.probability = value;
|
self.probability = value;
|
||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 设置颜色过滤约束。
|
||||||
pub fn color_filter<T>(mut self, filter: T) -> Self
|
pub fn color_filter<T>(mut self, filter: T) -> Self
|
||||||
where
|
where
|
||||||
T: ColorFilter + Send + Sync + 'static,
|
T: ColorFilter + Send + Sync + 'static,
|
||||||
@@ -45,6 +49,7 @@ impl OcrBuilder {
|
|||||||
self
|
self
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 设置字符集限制。
|
||||||
pub fn charset_restrict<T>(mut self, restrict: T) -> Self
|
pub fn charset_restrict<T>(mut self, restrict: T) -> Self
|
||||||
where
|
where
|
||||||
T: TokenFilter + Send + Sync + 'static,
|
T: TokenFilter + Send + Sync + 'static,
|
||||||
@@ -66,7 +71,6 @@ impl OcrBuilder {
|
|||||||
None => None,
|
None => None,
|
||||||
};
|
};
|
||||||
|
|
||||||
// Ocr::new(session, self)
|
|
||||||
Ocr {
|
Ocr {
|
||||||
runtime,
|
runtime,
|
||||||
png_fix: self.png_fix, // 原地解构出来
|
png_fix: self.png_fix, // 原地解构出来
|
||||||
|
|||||||
@@ -9,22 +9,22 @@ use std::collections::HashMap;
|
|||||||
/// 字符集:token 列表与索引的双向映射。
|
/// 字符集:token 列表与索引的双向映射。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct Charset {
|
pub struct Charset {
|
||||||
|
/// 字符集 token 列表。
|
||||||
// 使用 Cow 统一静态切片和动态读取的 Vec<String>,内部实现真正的零拷贝
|
// 使用 Cow 统一静态切片和动态读取的 Vec<String>,内部实现真正的零拷贝
|
||||||
pub tokens: Vec<Cow<'static, str>>,
|
pub tokens: Vec<Cow<'static, str>>,
|
||||||
|
/// 字符到索引的反查表。
|
||||||
// 反向查找表,保证字符转索引为 O(1)
|
// 反向查找表,保证字符转索引为 O(1)
|
||||||
pub char_to_idx: HashMap<Cow<'static, str>, usize>,
|
pub char_to_idx: HashMap<Cow<'static, str>, usize>,
|
||||||
// 当前处于激活状态的有效索引缓存 (用于 CTC 解码前的过滤加速)
|
|
||||||
// pub valid_indices: HashSet<usize>,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Charset {
|
impl Charset {
|
||||||
// 内部底层统一收拢构造
|
// 内部底层统一收拢构造
|
||||||
|
/// 从 token 列表构建字符集。
|
||||||
pub fn new(tokens: Vec<Cow<'static, str>>) -> Self {
|
pub fn new(tokens: Vec<Cow<'static, str>>) -> Self {
|
||||||
let mut char_to_idx = HashMap::with_capacity(tokens.len());
|
let mut char_to_idx = HashMap::with_capacity(tokens.len());
|
||||||
for (idx, token) in tokens.iter().enumerate() {
|
for (idx, token) in tokens.iter().enumerate() {
|
||||||
char_to_idx.entry(token.clone()).or_insert(idx);
|
char_to_idx.entry(token.clone()).or_insert(idx);
|
||||||
// 如果字符集有重复,保留第一个遇到的索引 (符合 Python .index 逻辑)
|
// 如果字符集有重复,保留第一个遇到的索引 (符合 Python .index 逻辑)
|
||||||
// char_to_idx.entry(token.to_string()).or_insert(idx);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
Self {
|
Self {
|
||||||
@@ -49,9 +49,11 @@ impl Charset {
|
|||||||
self.tokens.get(index).map(|cow| cow.as_ref())
|
self.tokens.get(index).map(|cow| cow.as_ref())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// 判断字符是否在字符集中。
|
||||||
pub fn is_valid_char(&self, char_str: &str) -> bool {
|
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 {
|
pub fn size(&self) -> usize {
|
||||||
self.tokens.len()
|
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 容器
|
// 3. 将扁平字节数组重新打包回 DynamicImage 容器
|
||||||
let filtered_buffer = ImageBuffer::<Rgb<u8>, Vec<u8>>::from_raw(width, height, raw_pixels)
|
let filtered_buffer = ImageBuffer::<Rgb<u8>, Vec<u8>>::from_raw(width, height, raw_pixels)
|
||||||
// .ok_or_else(|| anyhow!("图像缓冲重新组装失败,维度与数据大小不匹配"))?;
|
.ok_or(ImagePreprocessError::BufferLengthMismatch {
|
||||||
.ok_or_else(|| ImagePreprocessError::BufferLengthMismatch {
|
|
||||||
expected: expected_len,
|
expected: expected_len,
|
||||||
actual: actual_len,
|
actual: actual_len,
|
||||||
width,
|
width,
|
||||||
@@ -62,11 +61,14 @@ pub fn apply_to_image(
|
|||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
||||||
/// HSV 颜色区间,下界与上界各为 `(H, S, V)`。
|
/// HSV 颜色区间,下界与上界各为 `(H, S, V)`。
|
||||||
pub struct HsvRange {
|
pub struct HsvRange {
|
||||||
|
/// 区间下界 `(H, S, V)`。
|
||||||
pub lower: (u8, u8, u8), // (H, S, V)
|
pub lower: (u8, u8, u8), // (H, S, V)
|
||||||
|
/// 区间上界 `(H, S, V)`。
|
||||||
pub upper: (u8, u8, u8), // (H, S, V)
|
pub upper: (u8, u8, u8), // (H, S, V)
|
||||||
}
|
}
|
||||||
|
|
||||||
impl HsvRange {
|
impl HsvRange {
|
||||||
|
/// 创建 HSV 区间。
|
||||||
pub const fn new(lower: (u8, u8, u8), upper: (u8, u8, u8)) -> Self {
|
pub const fn new(lower: (u8, u8, u8), upper: (u8, u8, u8)) -> Self {
|
||||||
Self { lower, upper }
|
Self { lower, upper }
|
||||||
}
|
}
|
||||||
@@ -76,7 +78,6 @@ impl HsvRange {
|
|||||||
pub fn validate(&self) -> Result<(), ImagePreprocessError> {
|
pub fn validate(&self) -> Result<(), ImagePreprocessError> {
|
||||||
// 1. 校验 H 通道边界 (OpenCV 中 H 范围是 0-180)
|
// 1. 校验 H 通道边界 (OpenCV 中 H 范围是 0-180)
|
||||||
if self.lower.0 > 180 || self.upper.0 > 180 {
|
if self.lower.0 > 180 || self.upper.0 > 180 {
|
||||||
// return Err("H通道值必须在 0-180 范围内".to_string());
|
|
||||||
return Err(ImagePreprocessError::InvalidHsvRange(
|
return Err(ImagePreprocessError::InvalidHsvRange(
|
||||||
"H通道值必须在 0-180 范围内".to_string(),
|
"H通道值必须在 0-180 范围内".to_string(),
|
||||||
));
|
));
|
||||||
@@ -85,7 +86,6 @@ impl HsvRange {
|
|||||||
// 2. 校验下界不能大于上界
|
// 2. 校验下界不能大于上界
|
||||||
if self.lower.0 > self.upper.0 || self.lower.1 > self.upper.1 || self.lower.2 > self.upper.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(
|
return Err(ImagePreprocessError::InvalidHsvRange(
|
||||||
"HSV范围下界不能大于上界".to_string(),
|
"HSV范围下界不能大于上界".to_string(),
|
||||||
));
|
));
|
||||||
@@ -97,16 +97,27 @@ impl HsvRange {
|
|||||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||||
/// 颜色预设:常见颜色对应的 HSV 区间集合。
|
/// 颜色预设:常见颜色对应的 HSV 区间集合。
|
||||||
pub enum ColorPreset {
|
pub enum ColorPreset {
|
||||||
|
/// 红色。
|
||||||
Red,
|
Red,
|
||||||
|
/// 蓝色。
|
||||||
Blue,
|
Blue,
|
||||||
|
/// 绿色。
|
||||||
Green,
|
Green,
|
||||||
|
/// 黄色。
|
||||||
Yellow,
|
Yellow,
|
||||||
|
/// 橙色。
|
||||||
Orange,
|
Orange,
|
||||||
|
/// 紫色。
|
||||||
Purple,
|
Purple,
|
||||||
|
/// 青色。
|
||||||
Cyan,
|
Cyan,
|
||||||
|
/// 黑色。
|
||||||
Black,
|
Black,
|
||||||
|
/// 白色。
|
||||||
White,
|
White,
|
||||||
|
/// 灰色。
|
||||||
Gray,
|
Gray,
|
||||||
|
/// 自定义区间列表。
|
||||||
Custom(Vec<HsvRange>),
|
Custom(Vec<HsvRange>),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -193,7 +204,6 @@ impl FromStr for ColorPreset {
|
|||||||
"black" => Ok(ColorPreset::Black),
|
"black" => Ok(ColorPreset::Black),
|
||||||
"white" => Ok(ColorPreset::White),
|
"white" => Ok(ColorPreset::White),
|
||||||
"gray" => Ok(ColorPreset::Gray),
|
"gray" => Ok(ColorPreset::Gray),
|
||||||
// _ => Err(format!("不支持的颜色预设: {}", s)),
|
|
||||||
_ => Err(ImagePreprocessError::UnknownColorPreset(s.to_string())),
|
_ => Err(ImagePreprocessError::UnknownColorPreset(s.to_string())),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -205,12 +215,15 @@ impl FromStr for ColorPreset {
|
|||||||
|
|
||||||
/// 颜色匹配上下文:当前像素的 HSV 值。
|
/// 颜色匹配上下文:当前像素的 HSV 值。
|
||||||
pub struct PixelCtx {
|
pub struct PixelCtx {
|
||||||
|
/// 当前像素的 HSV 值。
|
||||||
pub hsv: (u8, u8, u8),
|
pub hsv: (u8, u8, u8),
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 颜色过滤约束接口:提供一组 HSV 区间。
|
/// 颜色过滤约束接口:提供一组 HSV 区间。
|
||||||
pub trait ColorFilter {
|
pub trait ColorFilter {
|
||||||
|
/// 将有效区间追加到目标容器。
|
||||||
fn append_ranges(&self, target: &mut Vec<HsvRange>);
|
fn append_ranges(&self, target: &mut Vec<HsvRange>);
|
||||||
|
/// 预估有效区间数量。
|
||||||
fn estimated_count(&self) -> usize;
|
fn estimated_count(&self) -> usize;
|
||||||
/// 验证过滤器配置是否合法,默认直接放行。
|
/// 验证过滤器配置是否合法,默认直接放行。
|
||||||
fn validate_self(&self) -> Result<(), ImagePreprocessError> {
|
fn validate_self(&self) -> Result<(), ImagePreprocessError> {
|
||||||
@@ -258,6 +271,7 @@ impl ColorFilter for ColorPreset {
|
|||||||
|
|
||||||
/// 多路颜色“或”逻辑组合子(并集网络)
|
/// 多路颜色“或”逻辑组合子(并集网络)
|
||||||
pub struct MultiOrColorRestrict<'a> {
|
pub struct MultiOrColorRestrict<'a> {
|
||||||
|
/// 参与「或」组合的过滤器列表。
|
||||||
pub filters: Vec<&'a dyn ColorFilter>,
|
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::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::error::{ImagePreprocessError, Result, TensorError};
|
||||||
|
use crate::ocr::color_filter::{HsvRange, apply_to_image};
|
||||||
use crate::traits::OcrEngine;
|
use crate::traits::OcrEngine;
|
||||||
use crate::utils::image_convert::png_rgba_white_preprocess;
|
use crate::utils::image_convert::png_rgba_white_preprocess;
|
||||||
use crate::utils::image_processor::{convert_to_grayscale, resize_image};
|
use crate::utils::image_processor::{convert_to_grayscale, resize_image};
|
||||||
@@ -13,7 +12,7 @@ use image::DynamicImage;
|
|||||||
use ndarray::ArrayView2;
|
use ndarray::ArrayView2;
|
||||||
use std::borrow::Cow;
|
use std::borrow::Cow;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
use tracing::warn;
|
use tracing::{debug, warn};
|
||||||
/// OCR 识别结果:纯文本或携带概率的文本。
|
/// OCR 识别结果:纯文本或携带概率的文本。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub enum OcrResult {
|
pub enum OcrResult {
|
||||||
@@ -21,6 +20,7 @@ pub enum OcrResult {
|
|||||||
Text(String),
|
Text(String),
|
||||||
/// 携带概率的结果(`probability = true` 时返回)。
|
/// 携带概率的结果(`probability = true` 时返回)。
|
||||||
Probability {
|
Probability {
|
||||||
|
/// 识别出的文本。
|
||||||
text: String,
|
text: String,
|
||||||
/// 全量概率矩阵 `[Steps, Classes]`。
|
/// 全量概率矩阵 `[Steps, Classes]`。
|
||||||
probabilities: Vec<Vec<f32>>,
|
probabilities: Vec<Vec<f32>>,
|
||||||
@@ -28,7 +28,10 @@ pub enum OcrResult {
|
|||||||
confidence: f64,
|
confidence: f64,
|
||||||
},
|
},
|
||||||
/// 不支持的模型或未知输出。
|
/// 不支持的模型或未知输出。
|
||||||
Unsupported { message: String },
|
Unsupported {
|
||||||
|
/// 不支持原因说明。
|
||||||
|
message: String,
|
||||||
|
},
|
||||||
}
|
}
|
||||||
impl OcrResult {
|
impl OcrResult {
|
||||||
/// 消费自身并提取最终文本。
|
/// 消费自身并提取最终文本。
|
||||||
@@ -114,6 +117,7 @@ pub struct Ocr<'a> {
|
|||||||
impl<'a> Ocr<'a> {
|
impl<'a> Ocr<'a> {
|
||||||
// 初始化任务,设置默认参数
|
// 初始化任务,设置默认参数
|
||||||
|
|
||||||
|
/// 绑定引擎会话创建 OCR 识别器。
|
||||||
pub fn new(runtime: &'a dyn OcrEngine) -> Self {
|
pub fn new(runtime: &'a dyn OcrEngine) -> Self {
|
||||||
Ocr {
|
Ocr {
|
||||||
runtime,
|
runtime,
|
||||||
@@ -123,6 +127,7 @@ impl<'a> Ocr<'a> {
|
|||||||
final_charset_indices: None,
|
final_charset_indices: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
/// 创建 OCR 构建器。
|
||||||
pub fn builder() -> OcrBuilder {
|
pub fn builder() -> OcrBuilder {
|
||||||
OcrBuilder::default()
|
OcrBuilder::default()
|
||||||
}
|
}
|
||||||
@@ -130,7 +135,7 @@ impl<'a> Ocr<'a> {
|
|||||||
impl<'a> Ocr<'a> {
|
impl<'a> Ocr<'a> {
|
||||||
/// 对输入图像执行 OCR 识别并返回结果。
|
/// 对输入图像执行 OCR 识别并返回结果。
|
||||||
pub fn predict(&self, image: &DynamicImage) -> Result<OcrResult> {
|
pub fn predict(&self, image: &DynamicImage) -> Result<OcrResult> {
|
||||||
println!("当前颜色过滤器状态: {:?}", self.final_color_ranges);
|
debug!("当前颜色过滤器状态: {:?}", self.final_color_ranges);
|
||||||
|
|
||||||
// =====================================================================
|
// =====================================================================
|
||||||
// 管道节点 1: 颜色过滤流水线
|
// 管道节点 1: 颜色过滤流水线
|
||||||
@@ -139,10 +144,6 @@ impl<'a> Ocr<'a> {
|
|||||||
// =====================================================================
|
// =====================================================================
|
||||||
let img_cow = match &self.final_color_ranges {
|
let img_cow = match &self.final_color_ranges {
|
||||||
Err(err_msg) => {
|
Err(err_msg) => {
|
||||||
// return Err(anyhow::anyhow!(
|
|
||||||
// "颜色过滤器初始化失败,全链路短路: {}",
|
|
||||||
// err_msg
|
|
||||||
// ));
|
|
||||||
return Err(ImagePreprocessError::FilterConfigInvalid(
|
return Err(ImagePreprocessError::FilterConfigInvalid(
|
||||||
err_msg.to_string(),
|
err_msg.to_string(),
|
||||||
))?;
|
))?;
|
||||||
@@ -162,17 +163,6 @@ impl<'a> Ocr<'a> {
|
|||||||
let raw_tensor = self.runtime.inference(tensor)?;
|
let raw_tensor = self.runtime.inference(tensor)?;
|
||||||
|
|
||||||
// 3. 后处理分流:直接返回 OcrResult
|
// 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)?;
|
let ocr_output = self.process_model_output(raw_tensor)?;
|
||||||
Ok(ocr_output)
|
Ok(ocr_output)
|
||||||
}
|
}
|
||||||
@@ -216,7 +206,7 @@ impl<'a> Ocr<'a> {
|
|||||||
1 => {
|
1 => {
|
||||||
let gray_img = convert_to_grayscale(&resized_img);
|
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),
|
(1, 1, target_h as usize, target_w as usize),
|
||||||
|(_, _, y, x)| {
|
|(_, _, y, x)| {
|
||||||
let pixel = gray_img.get_pixel(x as u32, y as u32)[0] as f32;
|
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
|
// (pixel / 255.0 - 0.5) / 0.5
|
||||||
norm.normalize(pixel)
|
norm.normalize(pixel)
|
||||||
},
|
},
|
||||||
);
|
)
|
||||||
array
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- 情况 B: 三通道(RGB),对应 Python 的 transpose(2, 0, 1) 的 CHW 布局 ---
|
// --- 情况 B: 三通道(RGB),对应 Python 的 transpose(2, 0, 1) 的 CHW 布局 ---
|
||||||
3 => {
|
3 => {
|
||||||
let rgb_img = resized_img.to_rgb8();
|
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),
|
(1, 3, target_h as usize, target_w as usize),
|
||||||
|(_, c, y, x)| {
|
|(_, c, y, x)| {
|
||||||
let pixel = rgb_img.get_pixel(x as u32, y as u32)[c] as f32;
|
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
|
// (pixel / 255.0 - 0.5) / 0.5
|
||||||
norm.normalize(pixel)
|
norm.normalize(pixel)
|
||||||
},
|
},
|
||||||
);
|
)
|
||||||
// Tensor::from(array)
|
|
||||||
array
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// _ => return Err(anyhow::anyhow!("不支持的通道数配置: {}", meta.channel)),
|
// _ => return Err(anyhow::anyhow!("不支持的通道数配置: {}", meta.channel)),
|
||||||
@@ -255,15 +242,11 @@ impl<'a> Ocr<'a> {
|
|||||||
Ok(array4)
|
Ok(array4)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 这段代码未来直接放入 ddddocr-core
|
|
||||||
fn process_model_output(&self, output: OcrOutput) -> Result<OcrResult, TensorError> {
|
fn process_model_output(&self, output: OcrOutput) -> Result<OcrResult, TensorError> {
|
||||||
match output {
|
match output {
|
||||||
OcrOutput::Indices(array1) => {
|
OcrOutput::Indices(array1) => {
|
||||||
// 对应你原来的 process_i64_tensor
|
// 对应原来的 process_i64_tensor
|
||||||
let slice = array1
|
let slice = array1.as_slice().ok_or(TensorError::NonContiguousMemory)?;
|
||||||
.as_slice()
|
|
||||||
// .ok_or_else(|| anyhow::anyhow!("内存不连续,无法执行零拷贝解码"))?;
|
|
||||||
.ok_or_else(|| TensorError::NonContiguousMemory)?;
|
|
||||||
let final_text = self.ctc_decode_to_string(slice);
|
let final_text = self.ctc_decode_to_string(slice);
|
||||||
|
|
||||||
if self.probability {
|
if self.probability {
|
||||||
@@ -277,7 +260,7 @@ impl<'a> Ocr<'a> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
OcrOutput::Logits(matrix_view) => {
|
OcrOutput::Logits(matrix_view) => {
|
||||||
// 对应你原来的 process_f32_tensor
|
// 对应原来的 process_f32_tensor
|
||||||
// 注意:此时的 matrix_view 已经是干净的标准的 ndarray::Array2<f32>,且保证是 [Steps, Classes] 2D 形状
|
// 注意:此时的 matrix_view 已经是干净的标准的 ndarray::Array2<f32>,且保证是 [Steps, Classes] 2D 形状
|
||||||
if self.probability {
|
if self.probability {
|
||||||
let (probabilities_list, confidence, predicted_indices) =
|
let (probabilities_list, confidence, predicted_indices) =
|
||||||
@@ -332,6 +315,7 @@ impl<'a> Ocr<'a> {
|
|||||||
None => tokens.iter().map(|cow| cow.as_ref()).collect(),
|
None => tokens.iter().map(|cow| cow.as_ref()).collect(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
/// 返回当前生效的可用 token 数量。
|
||||||
pub fn valid_size(&self) -> usize {
|
pub fn valid_size(&self) -> usize {
|
||||||
match &self.final_charset_indices {
|
match &self.final_charset_indices {
|
||||||
Some(indices) => indices.len(),
|
Some(indices) => indices.len(),
|
||||||
@@ -390,100 +374,10 @@ impl<'a> Ocr<'a> {
|
|||||||
|
|
||||||
(probabilities_list, confidence, predicted_indices)
|
(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 {
|
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 charset = &self.runtime.metadata().charset;
|
||||||
let tokens = &charset.tokens;
|
let tokens = &charset.tokens;
|
||||||
// let valid_indices = &charset.valid_indices;
|
|
||||||
|
|
||||||
// 对应 _ctc_decode_indices 的逻辑:去重、去 blank (0)
|
// 对应 _ctc_decode_indices 的逻辑:去重、去 blank (0)
|
||||||
let mut res = String::new();
|
let mut res = String::new();
|
||||||
@@ -509,10 +403,10 @@ impl<'a> Ocr<'a> {
|
|||||||
|
|
||||||
// 史诗级加速点:如果是 None,说明没限制,根本不进入分支,直接放行!
|
// 史诗级加速点:如果是 None,说明没限制,根本不进入分支,直接放行!
|
||||||
// 只有当有具体限制(Some)时,才去跑 4-5 次 CPU 寄存器级别的二分查找
|
// 只有当有具体限制(Some)时,才去跑 4-5 次 CPU 寄存器级别的二分查找
|
||||||
if let Some(ref indices) = self.final_charset_indices {
|
if let Some(ref indices) = self.final_charset_indices
|
||||||
if indices.binary_search(&u_idx).is_err() {
|
&& indices.binary_search(&u_idx).is_err()
|
||||||
continue;
|
{
|
||||||
}
|
continue;
|
||||||
}
|
}
|
||||||
|
|
||||||
// 5. 字符映射
|
// 5. 字符映射
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ pub enum Normalization {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl Normalization {
|
impl Normalization {
|
||||||
|
/// 对像素值执行归一化。
|
||||||
#[inline(always)]
|
#[inline(always)]
|
||||||
pub fn normalize(&self, pixel: f32) -> f32 {
|
pub fn normalize(&self, pixel: f32) -> f32 {
|
||||||
match self {
|
match self {
|
||||||
@@ -40,16 +41,20 @@ pub enum Resize {
|
|||||||
/// OCR 模型元数据:字符集、缩放策略、通道数与归一化配置。
|
/// OCR 模型元数据:字符集、缩放策略、通道数与归一化配置。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct ModelMetadata {
|
pub struct ModelMetadata {
|
||||||
|
/// 字符集。
|
||||||
pub charset: Charset,
|
pub charset: Charset,
|
||||||
/// 是否为单字识别模型
|
/// 是否为单字识别模型
|
||||||
pub word: bool,
|
pub word: bool,
|
||||||
|
/// 缩放策略。
|
||||||
pub resize: Resize,
|
pub resize: Resize,
|
||||||
/// 图像通道数(1 或 3)
|
/// 图像通道数(1 或 3)
|
||||||
pub channel: u8,
|
pub channel: u8,
|
||||||
|
/// 像素归一化配置。
|
||||||
pub normalization: Normalization,
|
pub normalization: Normalization,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ModelMetadata {
|
impl ModelMetadata {
|
||||||
|
/// 创建模型元数据。
|
||||||
pub fn new(
|
pub fn new(
|
||||||
charset: Charset,
|
charset: Charset,
|
||||||
word: bool,
|
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,21 +1,25 @@
|
|||||||
//! 字符集限制:按字符属性或索引过滤识别范围。
|
//! 字符集限制:按字符属性或索引过滤识别范围。
|
||||||
|
|
||||||
use std::borrow::Cow;
|
use std::borrow::Cow;
|
||||||
|
use tracing::warn;
|
||||||
|
|
||||||
/// 字符集校验上下文:当前 token 的文本与索引。
|
/// 字符集校验上下文:当前 token 的文本与索引。
|
||||||
pub struct ValidationCtx<'a> {
|
pub struct ValidationCtx<'a> {
|
||||||
pub text: &'a str, // 当前 Token 的文本内容
|
/// 当前 token 的文本内容。
|
||||||
|
pub text: &'a str, // 当前 Token 的文本内容
|
||||||
|
/// 当前 token 的 ID 索引。
|
||||||
pub token_id: usize, // 当前 Token 的 ID 索引
|
pub token_id: usize, // 当前 Token 的 ID 索引
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 字符集限制接口:决定某个 token 是否放行。
|
/// 字符集限制接口:决定某个 token 是否放行。
|
||||||
pub trait TokenFilter {
|
pub trait TokenFilter {
|
||||||
|
/// 判断 token 是否放行。
|
||||||
fn matches(&self, ctx: &ValidationCtx) -> bool;
|
fn matches(&self, ctx: &ValidationCtx) -> bool;
|
||||||
/// 预估匹配数量的容量提示。
|
/// 预估匹配数量的容量提示。
|
||||||
fn estimated_capacity(&self) -> usize {
|
fn estimated_capacity(&self) -> usize {
|
||||||
128
|
128
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 遍历全量字符集筛选可用索引(放行 CTC blank、排序去重、空交集返回 `None`)。
|
/// 遍历全量字符集筛选可用索引(放行 CTC blank、排序去重、空交集返回 `None`)。
|
||||||
fn apply_to_charset(&self, tokens: &[Cow<str>]) -> Option<Vec<usize>> {
|
fn apply_to_charset(&self, tokens: &[Cow<str>]) -> Option<Vec<usize>> {
|
||||||
let mut has_any_match = false;
|
let mut has_any_match = false;
|
||||||
@@ -49,7 +53,7 @@ pub trait TokenFilter {
|
|||||||
|
|
||||||
// 3. 终极防御:如果整个模型字符集除了 Blank,一个都没对上,直接退化为 None(全量识别)
|
// 3. 终极防御:如果整个模型字符集除了 Blank,一个都没对上,直接退化为 None(全量识别)
|
||||||
if !has_any_match {
|
if !has_any_match {
|
||||||
println!("警告:当前限制策略与模型字符集完全没有交集!已自动恢复全量识别。");
|
warn!("当前限制策略与模型字符集完全没有交集,已自动恢复全量识别");
|
||||||
None
|
None
|
||||||
} else {
|
} else {
|
||||||
// 4. 排序并去重,为 Ocr 引擎后续进行极其高频的『二分查找』筑起绝对安全的底层保障
|
// 4. 排序并去重,为 Ocr 引擎后续进行极其高频的『二分查找』筑起绝对安全的底层保障
|
||||||
@@ -63,9 +67,13 @@ pub trait TokenFilter {
|
|||||||
/// 按字符属性限制:数字、大小写字母或自定义列表。
|
/// 按字符属性限制:数字、大小写字母或自定义列表。
|
||||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||||
pub enum CharRestrict {
|
pub enum CharRestrict {
|
||||||
|
/// 仅数字。
|
||||||
Digit,
|
Digit,
|
||||||
|
/// 仅小写字母。
|
||||||
Lowercase,
|
Lowercase,
|
||||||
|
/// 仅大写字母。
|
||||||
Uppercase,
|
Uppercase,
|
||||||
|
/// 自定义字符列表。
|
||||||
CustomList(Vec<String>),
|
CustomList(Vec<String>),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -90,8 +98,11 @@ impl TokenFilter for CharRestrict {
|
|||||||
/// 按索引限制:前 N 个、索引范围或索引列表。
|
/// 按索引限制:前 N 个、索引范围或索引列表。
|
||||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||||
pub enum IdRestrict {
|
pub enum IdRestrict {
|
||||||
|
/// 前 N 个索引。
|
||||||
TopN(usize),
|
TopN(usize),
|
||||||
|
/// 指定索引范围。
|
||||||
IdRange(std::ops::Range<usize>),
|
IdRange(std::ops::Range<usize>),
|
||||||
|
/// 指定索引列表。
|
||||||
IdList(Vec<usize>),
|
IdList(Vec<usize>),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,6 +128,7 @@ impl TokenFilter for IdRestrict {
|
|||||||
|
|
||||||
/// 多路“或”逻辑组合子(支持 N 个规则无缝并集)
|
/// 多路“或”逻辑组合子(支持 N 个规则无缝并集)
|
||||||
pub struct MultiOrRestrict<'a> {
|
pub struct MultiOrRestrict<'a> {
|
||||||
|
/// 参与「或」组合的过滤器列表。
|
||||||
pub filters: Vec<&'a dyn TokenFilter>,
|
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 imageproc::template_matching::{MatchTemplateMethod, match_template};
|
||||||
use ndarray::{ArrayView2, ArrayView3};
|
use ndarray::{ArrayView2, ArrayView3};
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
use tracing::debug;
|
||||||
|
|
||||||
/// 滑块匹配结果:检测中心坐标与置信度。
|
/// 滑块匹配结果:检测中心坐标与置信度。
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub struct SlideResult {
|
pub struct SlideResult {
|
||||||
/// 检测中心坐标 `[x, y]`。
|
/// 检测中心坐标 `[x, y]`。
|
||||||
pub target: [i32; 2],
|
pub target: [i32; 2],
|
||||||
|
/// 检测中心的 x 坐标。
|
||||||
pub target_x: i32,
|
pub target_x: i32,
|
||||||
|
/// 检测中心的 y 坐标。
|
||||||
pub target_y: i32,
|
pub target_y: i32,
|
||||||
|
/// 匹配置信度。
|
||||||
pub confidence: f64,
|
pub confidence: f64,
|
||||||
}
|
}
|
||||||
impl fmt::Display for SlideResult {
|
impl fmt::Display for SlideResult {
|
||||||
@@ -35,9 +39,11 @@ impl fmt::Display for SlideResult {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 滑块匹配服务:提供模板匹配与差异比较两种识别模式。
|
/// 滑块匹配服务:提供模板匹配与差异比较两种识别模式。
|
||||||
|
#[derive(Default)]
|
||||||
pub struct Slider;
|
pub struct Slider;
|
||||||
|
|
||||||
impl Slider {
|
impl Slider {
|
||||||
|
/// 创建滑块匹配服务。
|
||||||
pub fn new() -> Self {
|
pub fn new() -> Self {
|
||||||
Self
|
Self
|
||||||
}
|
}
|
||||||
@@ -113,7 +119,6 @@ impl Slider {
|
|||||||
let background_label = Luma([0u8]);
|
let background_label = Luma([0u8]);
|
||||||
let labelled = connected_components(&cleaned, Connectivity::Eight, background_label);
|
let labelled = connected_components(&cleaned, Connectivity::Eight, background_label);
|
||||||
|
|
||||||
// // 统计每个标签出现的频率(即面积)
|
|
||||||
// 4. 寻找最大连通区域 (对应 findContours + max area)
|
// 4. 寻找最大连通区域 (对应 findContours + max area)
|
||||||
if let Some(max_label) = image_processor::find_contours_and_max(&labelled) {
|
if let Some(max_label) = image_processor::find_contours_and_max(&labelled) {
|
||||||
// 5. 计算最大区域的边界框 (对应 cv2.boundingRect)
|
// 5. 计算最大区域的边界框 (对应 cv2.boundingRect)
|
||||||
@@ -139,7 +144,6 @@ impl Slider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// 模板匹配核心实现(对应 Python 的 `_perform_slide_match`)。
|
/// 模板匹配核心实现(对应 Python 的 `_perform_slide_match`)。
|
||||||
// 在 SlideEngine 中修改此入口进行测试
|
|
||||||
fn perform_slide_match(
|
fn perform_slide_match(
|
||||||
&self,
|
&self,
|
||||||
target: ArrayView3<u8>,
|
target: ArrayView3<u8>,
|
||||||
@@ -158,7 +162,6 @@ impl Slider {
|
|||||||
}
|
}
|
||||||
if th > bh || tw > bw {
|
if th > bh || tw > bw {
|
||||||
return Err(ImagePreprocessError::TargetExceedsBackground {
|
return Err(ImagePreprocessError::TargetExceedsBackground {
|
||||||
// "尺寸不匹配:滑块模板(target)尺寸 [{}x{}] 不能大于背景图(background) [{}x{}]",
|
|
||||||
target_w: tw,
|
target_w: tw,
|
||||||
target_h: th,
|
target_h: th,
|
||||||
bg_w: bw,
|
bg_w: bw,
|
||||||
@@ -195,7 +198,6 @@ impl Slider {
|
|||||||
// 转换逻辑 (假设你已经有方法转回 ImageBuffer)
|
// 转换逻辑 (假设你已经有方法转回 ImageBuffer)
|
||||||
let t_buf = ndarray_to_luma8(target);
|
let t_buf = ndarray_to_luma8(target);
|
||||||
let b_buf = ndarray_to_luma8(background);
|
let b_buf = ndarray_to_luma8(background);
|
||||||
// t_buf.save("debug_rust_target.png").unwrap();
|
|
||||||
|
|
||||||
// 2. 调用 imageproc 的 NCC 算法 (等价于 cv2.TM_CCOEFF_NORMED)
|
// 2. 调用 imageproc 的 NCC 算法 (等价于 cv2.TM_CCOEFF_NORMED)
|
||||||
// 模板匹配 (完全对齐 cv2.matchTemplate(..., cv2.TM_CCOEFF_NORMED))
|
// 模板匹配 (完全对齐 cv2.matchTemplate(..., cv2.TM_CCOEFF_NORMED))
|
||||||
@@ -204,18 +206,13 @@ impl Slider {
|
|||||||
&t_buf,
|
&t_buf,
|
||||||
MatchTemplateMethod::CrossCorrelationNormalized,
|
MatchTemplateMethod::CrossCorrelationNormalized,
|
||||||
);
|
);
|
||||||
// save_rust_result(&result, "debug_rust_target2.png");
|
|
||||||
// 3. 寻找最大值 (等价于 cv2.minMaxLoc)
|
// 3. 寻找最大值 (等价于 cv2.minMaxLoc)
|
||||||
let (max_val, max_loc) = min_max_loc(&result);
|
let (max_val, max_loc) = min_max_loc(&result);
|
||||||
|
|
||||||
// 4. 计算中心点 (与 Python 逻辑完全一致)
|
// 4. 计算中心点 (与 Python 逻辑完全一致)
|
||||||
let (th, tw) = target.dim();
|
let (th, tw) = target.dim();
|
||||||
|
|
||||||
let (center_x, center_y) =
|
let (center_x, center_y) = image_processor::calculate_center(max_loc, tw, th);
|
||||||
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);
|
|
||||||
SlideResult {
|
SlideResult {
|
||||||
target: [center_x, center_y],
|
target: [center_x, center_y],
|
||||||
target_x: center_x,
|
target_x: center_x,
|
||||||
@@ -240,9 +237,6 @@ impl Slider {
|
|||||||
let target_edges = canny(&t_buf, 50.0, 150.0);
|
let target_edges = canny(&t_buf, 50.0, 150.0);
|
||||||
let background_edges = canny(&b_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))
|
// 3. 模板匹配 (完全对齐 cv2.matchTemplate(..., cv2.TM_CCOEFF_NORMED))
|
||||||
// 在边缘图上计算归一化互相关系数
|
// 在边缘图上计算归一化互相关系数
|
||||||
let result = match_template(
|
let result = match_template(
|
||||||
@@ -256,14 +250,12 @@ impl Slider {
|
|||||||
// 5. 计算中心位置 (对齐 Python 逻辑)
|
// 5. 计算中心位置 (对齐 Python 逻辑)
|
||||||
// target_w, target_h 来自输入数组的维度
|
// target_w, target_h 来自输入数组的维度
|
||||||
let (th, tw) = target.dim();
|
let (th, tw) = target.dim();
|
||||||
let (center_x, center_y) =
|
let (center_x, center_y) = image_processor::calculate_center(max_loc, tw, th);
|
||||||
image_processor::calculate_center(max_loc, tw as usize, th as usize);
|
|
||||||
|
|
||||||
// 打印调试信息,方便与 Python 对比
|
// 打印调试信息,方便与 Python 对比
|
||||||
// println!("Edge Match: max_val: {}, max_loc: {:?}", max_val, max_loc);
|
debug!("-Rust Target Width (tw): {}", tw);
|
||||||
println!("-Rust Target Width (tw): {}", tw);
|
debug!("-Rust Best Max Loc X: {}", max_loc.0);
|
||||||
println!("-Rust Best Max Loc X: {}", max_loc.0);
|
debug!("-Rust Final Center X: {}", center_x);
|
||||||
println!("-Rust Final Center X: {}", center_x);
|
|
||||||
SlideResult {
|
SlideResult {
|
||||||
target: [center_x, center_y],
|
target: [center_x, center_y],
|
||||||
target_x: center_x,
|
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 {
|
pub trait Info {
|
||||||
|
/// 获取输入张量信息列表。
|
||||||
fn input_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
|
fn input_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
|
||||||
|
/// 获取输出张量信息列表。
|
||||||
fn output_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
|
fn output_info(&self) -> crate::error::Result<Vec<TensorInfo>>;
|
||||||
|
/// 获取模型完整输入/输出信息。
|
||||||
fn model_info(&self) -> crate::error::Result<ModelInfo>;
|
fn model_info(&self) -> crate::error::Result<ModelInfo>;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -16,6 +19,7 @@ pub trait Info {
|
|||||||
pub trait InferenceEngine {
|
pub trait InferenceEngine {
|
||||||
/// 引擎产出的输出枚举(OCR 为 [`crate::OcrOutput`],检测为 [`crate::DetOutput`])。
|
/// 引擎产出的输出枚举(OCR 为 [`crate::OcrOutput`],检测为 [`crate::DetOutput`])。
|
||||||
type Output;
|
type Output;
|
||||||
|
/// 对输入张量执行推理并返回引擎定义的输出。
|
||||||
fn inference(
|
fn inference(
|
||||||
&self,
|
&self,
|
||||||
input_array: ndarray::Array4<f32>,
|
input_array: ndarray::Array4<f32>,
|
||||||
@@ -24,6 +28,7 @@ pub trait InferenceEngine {
|
|||||||
|
|
||||||
/// OCR 引擎接口:输出 [`crate::OcrOutput`],并提供模型元数据。
|
/// OCR 引擎接口:输出 [`crate::OcrOutput`],并提供模型元数据。
|
||||||
pub trait OcrEngine: InferenceEngine<Output = OcrOutput> + Info {
|
pub trait OcrEngine: InferenceEngine<Output = OcrOutput> + Info {
|
||||||
|
/// 获取模型元数据。
|
||||||
fn metadata(&self) -> &ModelMetadata;
|
fn metadata(&self) -> &ModelMetadata;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -32,13 +37,17 @@ pub trait DetEngine: InferenceEngine<Output = DetOutput> {}
|
|||||||
|
|
||||||
/// 模型加载器:从本地路径或字节流构建引擎会话。
|
/// 模型加载器:从本地路径或字节流构建引擎会话。
|
||||||
pub trait Loader {
|
pub trait Loader {
|
||||||
|
/// 构建出的引擎会话类型。
|
||||||
type Session;
|
type Session;
|
||||||
|
/// 构建过程中的错误类型。
|
||||||
type Error;
|
type Error;
|
||||||
|
/// 从本地模型路径构建会话。
|
||||||
fn build_for_path<P: AsRef<Path>>(
|
fn build_for_path<P: AsRef<Path>>(
|
||||||
&self,
|
&self,
|
||||||
model_path: P,
|
model_path: P,
|
||||||
) -> crate::error::Result<Self::Session, Self::Error>;
|
) -> crate::error::Result<Self::Session, Self::Error>;
|
||||||
|
|
||||||
|
/// 从模型字节流构建会话。
|
||||||
fn build_from_bytes(
|
fn build_from_bytes(
|
||||||
&self,
|
&self,
|
||||||
model_bytes: &[u8],
|
model_bytes: &[u8],
|
||||||
|
|||||||
@@ -3,19 +3,25 @@
|
|||||||
/// 张量元素的数据类型标记。
|
/// 张量元素的数据类型标记。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub enum TensorType {
|
pub enum TensorType {
|
||||||
|
/// 32 位浮点类型。
|
||||||
F32,
|
F32,
|
||||||
|
/// 64 位整数类型。
|
||||||
I64,
|
I64,
|
||||||
|
/// 其他类型。
|
||||||
Other,
|
Other,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 模型某个轴的维度特征:静态数值或动态符号。
|
/// 模型某个轴的维度特征:静态数值或动态符号。
|
||||||
#[derive(Clone, PartialEq, Eq)]
|
#[derive(Clone, PartialEq, Eq)]
|
||||||
pub enum AxisDim {
|
pub enum AxisDim {
|
||||||
|
/// 静态固定维度。
|
||||||
Static(usize),
|
Static(usize),
|
||||||
|
/// 动态符号维度。
|
||||||
Dynamic(String),
|
Dynamic(String),
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AxisDim {
|
impl AxisDim {
|
||||||
|
/// 是否为动态维度。
|
||||||
pub fn is_dynamic(&self) -> bool {
|
pub fn is_dynamic(&self) -> bool {
|
||||||
matches!(self, AxisDim::Dynamic(_))
|
matches!(self, AxisDim::Dynamic(_))
|
||||||
}
|
}
|
||||||
@@ -34,15 +40,20 @@ impl std::fmt::Debug for AxisDim {
|
|||||||
/// 单个张量(输入或输出)的名称、形状与数据类型描述。
|
/// 单个张量(输入或输出)的名称、形状与数据类型描述。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct TensorInfo {
|
pub struct TensorInfo {
|
||||||
|
/// 张量名称。
|
||||||
pub name: String,
|
pub name: String,
|
||||||
|
/// 各轴形状描述。
|
||||||
pub shape: Vec<AxisDim>,
|
pub shape: Vec<AxisDim>,
|
||||||
|
/// 元素数据类型。
|
||||||
pub tensor_type: TensorType,
|
pub tensor_type: TensorType,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 模型的完整输入/输出信息。
|
/// 模型的完整输入/输出信息。
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
pub struct ModelInfo {
|
pub struct ModelInfo {
|
||||||
|
/// 输入张量列表。
|
||||||
pub inputs: Vec<TensorInfo>,
|
pub inputs: Vec<TensorInfo>,
|
||||||
|
/// 输出张量列表。
|
||||||
pub outputs: Vec<TensorInfo>,
|
pub outputs: Vec<TensorInfo>,
|
||||||
/// 硬件执行提供者(`None` 表示使用引擎默认后端)。
|
/// 硬件执行提供者(`None` 表示使用引擎默认后端)。
|
||||||
pub providers: Option<Vec<String>>,
|
pub providers: Option<Vec<String>>,
|
||||||
|
|||||||
@@ -7,21 +7,22 @@ use ndarray::{Array3, ArrayViewD};
|
|||||||
/// 图像通道模式。
|
/// 图像通道模式。
|
||||||
#[derive(Debug)]
|
#[derive(Debug)]
|
||||||
pub enum ColorMode {
|
pub enum ColorMode {
|
||||||
|
/// RGB 三通道。
|
||||||
RGB,
|
RGB,
|
||||||
|
/// RGBA 四通道。
|
||||||
RGBA,
|
RGBA,
|
||||||
|
/// 灰度单通道。
|
||||||
L,
|
L,
|
||||||
}
|
}
|
||||||
|
|
||||||
/// 将 ndarray 数组转换为图像(自动识别 HWC 通道数)。
|
/// 将 ndarray 数组转换为图像(自动识别 HWC 通道数)。
|
||||||
// 对应 Python 版 _numpy_to_pil_image
|
// 对应 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 shape = array.shape();
|
||||||
let dim = shape.len();
|
let dim = shape.len();
|
||||||
|
|
||||||
// 1. 确保数据在内存中是连续的 (C order / Standard Layout)
|
// 1. 确保数据在内存中是连续的 (C order / Standard Layout)
|
||||||
// 如果 arr 是经过切片或转置的,这一步会进行必要的内存拷贝
|
// 如果 arr 是经过切片或转置的,这一步会进行必要的内存拷贝
|
||||||
// let standard = array.as_standard_layout();
|
|
||||||
// let (raw_data, _offset) = standard.to_owned().into_raw_vec_and_offset();
|
|
||||||
|
|
||||||
let color_mode = match dim {
|
let color_mode = match dim {
|
||||||
// 对应 Python: len(array.shape) == 2 (灰度图 H, W)
|
// 对应 Python: len(array.shape) == 2 (灰度图 H, W)
|
||||||
@@ -100,8 +101,10 @@ pub fn png_rgba_white_preprocess(img: &DynamicImage) -> DynamicImage {
|
|||||||
DynamicImage::ImageRgb8(background)
|
DynamicImage::ImageRgb8(background)
|
||||||
}
|
}
|
||||||
/// 将 DynamicImage 转换为 array 数组
|
/// 将 DynamicImage 转换为 array 数组
|
||||||
pub fn image_to_ndarray(image: &DynamicImage, mode: ColorMode) -> Result<Array3<u8>,ImagePreprocessError> {
|
pub fn image_to_ndarray(
|
||||||
// 1. 模式转换 (对应 utils.convert(target_mode)),此函数在时保留看后续优化是否需要替代image_to_ndarray
|
image: &DynamicImage,
|
||||||
|
mode: ColorMode,
|
||||||
|
) -> Result<Array3<u8>, ImagePreprocessError> {
|
||||||
// Rust utils 库通过 to_rgb8, to_luma8 等方法实现转换
|
// Rust utils 库通过 to_rgb8, to_luma8 等方法实现转换
|
||||||
let (width, height) = image.dimensions();
|
let (width, height) = image.dimensions();
|
||||||
|
|
||||||
@@ -116,7 +119,10 @@ pub fn image_to_ndarray(image: &DynamicImage, mode: ColorMode) -> Result<Array3<
|
|||||||
Ok(array)
|
Ok(array)
|
||||||
}
|
}
|
||||||
/// 将 array 数组转换为 DynamicImage
|
/// 将 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();
|
let shape = array.shape();
|
||||||
|
|
||||||
// 基础边界检查:至少要有 H 和 W 两个维度
|
// 基础边界检查:至少要有 H 和 W 两个维度
|
||||||
@@ -124,12 +130,15 @@ pub fn ndarray_to_image(array: ArrayViewD<u8>, mode: ColorMode) -> Result<Dynami
|
|||||||
return Err(ImagePreprocessError::InvalidDimensions {
|
return Err(ImagePreprocessError::InvalidDimensions {
|
||||||
expected: "至少为 2D array [H, W]".to_string(),
|
expected: "至少为 2D array [H, W]".to_string(),
|
||||||
actual: shape.to_vec(),
|
actual: shape.to_vec(),
|
||||||
})?;
|
});
|
||||||
}
|
}
|
||||||
from_ndarray(array, mode)
|
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();
|
let shape = array.shape();
|
||||||
|
|
||||||
// 映射:ndarray 的 shape 默认是 [Height, Width, (Channels)]
|
// 映射: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;
|
let expected_len = (width * height) as usize * channels;
|
||||||
|
|
||||||
// 构造通用错误闭包,避免 match 分支中重复编写冗长的错误对象
|
// 构造通用错误闭包,避免 match 分支中重复编写冗长的错误对象
|
||||||
let make_err = || {
|
let make_err = || ImagePreprocessError::BufferLengthMismatch {
|
||||||
ImagePreprocessError::BufferLengthMismatch {
|
expected: expected_len,
|
||||||
expected: expected_len,
|
actual: raw_len,
|
||||||
actual: raw_len,
|
width,
|
||||||
width,
|
height,
|
||||||
height,
|
channels,
|
||||||
channels,
|
|
||||||
}
|
|
||||||
};
|
};
|
||||||
|
|
||||||
// 2. 重新解释内存并构建 ImageBuffer
|
// 2. 重新解释内存并构建 ImageBuffer
|
||||||
@@ -175,3 +182,36 @@ fn from_ndarray(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,I
|
|||||||
.ok_or_else(make_err),
|
.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::error::{DdddError, Result};
|
||||||
use crate::utils::image_convert::ndarray_to_hwc_image;
|
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 image::DynamicImage;
|
||||||
use ndarray::ArrayViewD;
|
use ndarray::ArrayViewD;
|
||||||
use std::fmt;
|
use std::fmt;
|
||||||
|
|||||||
@@ -1,8 +1,8 @@
|
|||||||
//! 图像处理算法:OpenCV 风格的常用函数封装。
|
//! 图像处理算法: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};
|
use std::cmp::{max, min};
|
||||||
|
|
||||||
// 模拟openCV
|
// 模拟openCV
|
||||||
@@ -12,7 +12,7 @@ pub fn abs_diff(a: &ArrayView3<u8>, b: &ArrayView3<u8>) -> Array3<u8> {
|
|||||||
// 或者直接使用 zip_mut_with 处理以减少内存分配
|
// 或者直接使用 zip_mut_with 处理以减少内存分配
|
||||||
let mut diff = Array3::zeros(a.dim());
|
let mut diff = Array3::zeros(a.dim());
|
||||||
azip!((res in &mut diff, &va in a, &vb in b) {
|
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
|
diff
|
||||||
}
|
}
|
||||||
@@ -190,3 +190,43 @@ pub fn resize_image(
|
|||||||
image.resize_exact(target_width, target_height, FilterType::Lanczos3)
|
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;
|
use ndarray::s;
|
||||||
|
|
||||||
/// 将异构形状的模型输出规整为标准 `[Steps, Classes]` Logits 矩阵。
|
/// 将异构形状的模型输出规整为标准 `[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() {
|
let (steps, classes, data_dyn_view) = match shape.len() {
|
||||||
3 => {
|
3 => {
|
||||||
if shape[1] == 1 {
|
if shape[1] == 1 {
|
||||||
@@ -27,12 +30,10 @@ pub fn normalize_ocr_logits(array: ndarray::ArrayViewD<f32>, shape: &[usize]) ->
|
|||||||
// 我们把它虚构成一个 [1, Classes] 的 2D 矩阵来复用后面的 argmax 逻辑
|
// 我们把它虚构成一个 [1, Classes] 的 2D 矩阵来复用后面的 argmax 逻辑
|
||||||
1 => (1, shape[0], array),
|
1 => (1, shape[0], array),
|
||||||
_ => {
|
_ => {
|
||||||
return Err(
|
return Err(TensorError::DimensionMismatch {
|
||||||
TensorError::DimensionMismatch {
|
expected: "1D, 2D, or 3D OCR Logits".to_string(),
|
||||||
expected: "1D, 2D, or 3D OCR Logits".to_string(),
|
actual: shape.to_vec(),
|
||||||
actual: shape.to_vec(),
|
});
|
||||||
},
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -52,3 +53,32 @@ pub fn normalize_ocr_logits(array: ndarray::ArrayViewD<f32>, shape: &[usize]) ->
|
|||||||
|
|
||||||
Ok(OcrOutput::Logits(matrix_cow))
|
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