12 Commits

Author SHA1 Message Date
a3c4614574 refactor(core): 提炼公共类型
- 将 AxisDim、TensorInfo 等公共类型下沉至 ddddocr_core::types
- 项目结构优化
2026-07-30 16:55:58 +08:00
7d159c5702 feat(model): 新增 ModelLoader 链式构建 API 及 ORT GPU/Tract 多线程配置
- 在 ddddocr-core 中定义 ModelBuilder Trait 及其错误类型
- ddddocr-ort 支持 use_gpu、device_id 及 num_threads 链式配置与 CUDA 硬件加速
- ddddocr-tract 基于 multithread-mm 特性支持 CPU 线程数控制
- 支持基于 tract-linalg 配置推理线程数,显式引入 tract-linalg 的 multithread-mm 特性,控制 GEMM 算子并发
- 优化线程池加载策略,适配 Tokio 异步及 CLI 等多场景
2026-07-27 20:22:24 +08:00
44dae08221 refactor(load,tract):将 ModelMetadata JSON 加载逻辑解耦至 ddddocr-tract, 优化 Error 枚举结构与错误透传
- 在 load 模块中精简 Error 与 Result 别名定义
- 增加 ParseError 子类型区分路径与字节流加载失败
- 支持通过 #[from] 自动转换 Tract 引擎底层错误
- 移出 core 中的 serde 依赖,保持核心库纯洁
- 在 tract 中实现 TractModelMetadata 扩展 trait 加载解析配置
2026-07-23 13:43:32 +08:00
3499e89bf1 refactor(error): 规范化分层错误类型并优化异常捕捉
- 新增 tracing 记录异常,移除不必要的 Result
- 重构 错误处理架构
2026-07-20 20:18:38 +08:00
913ff4d884 refactor(errors): 重构错误处理,支持强类型匹配并剥离 base64 依赖
- 新增 Other变体以及构造函数new
- 剥离图像预处理中的 Base64 相关错误至业务层处理
- 引入强类型 `LogitsDimensionMismatch` 替代不便匹配的字符串错误
- 优化 `normalize_ocr_logits` 的转换流程,兼顾零拷贝性能与精细化报错
- 优化 全库错误处理
2026-07-17 20:08:32 +08:00
4f6987f594 refactor: 优化图像输入源设计,重构为零污染的 TryFromImage 特征
- 移除原有的 ImageInput 枚举,避免运行时匹配与所有权限制
- 引入自定义 TryFromImage 特征,专用于将不同来源安全转换为 ImageSource
- 优化错误处理新增 InvalidBase64Header 错误信息
- 迁移 load_image_from_input 到 image_helper.rs 为后续剥离到业务层做准备
2026-07-16 14:13:56 +08:00
cd70748393 refactor: 优化 ModelLoader 结构与维度解析逻辑
- 提取 `resolve_tensors` 公共函数,消除输入输出流的重复代码。
- 简化 `resolve_shape`,使其回归纯粹的维度翻译职责,移除过早的错误校验。
- 统一错误处理,将底层解析异常清晰映射至 `DdddError::InternalError`。
2026-07-14 18:56:41 +08:00
4fd38022fd refactor: 重构 core 包目录结构并消除旧版 mod.rs
- 优化 剥离 models,algo 层并平铺业务模块
- 重构 统一使用现代 filename.rs + 文件夹结构替代旧版 mod.rs
2026-07-11 17:38:30 +08:00
ea7fb43a14 refactor: 抽象解耦推理引擎并重构为多Crate工作空间架构
- 移除 核心层与 tract/Tensor 的强耦合,前/后处理全线转用标准 ndarray
- 针对 OCR 与目标检测(Det)分别设计独立的强类型输出小枚举(OcrOutput/DetOutput)
- 利用 Trait 关联类型(Associated Type)InferenceEngine,OcrEngine,DetEngine 统一接口,实现多后端解耦
- 引入 thiserror 库,建立完备的强类型错误处理机制(DdddError/Result)
- 完成项目结构初拆,剥离为 ddddocr-core 和 ddddocr-tract
2026-07-10 20:23:49 +08:00
2d9cb35590 feat(ocr,det,slide): 重构项目结构
- 优化 规范化模型目录
- 重构 Ocr,Detector,Slide 拆分规范化
2026-07-09 19:26:58 +08:00
0cf3d5fefb feat(ocr,det,slide): 重构配置解析流程,移除非必要的生命周期方法
- 优化 规范化模型目录
- 重构 Ocr,Detector配置解析流程
2026-07-08 15:48:56 +08:00
31271e80db refactor(slide,det): 优化项目结构,移除不必要的逻辑
- 优化 项目结构,移除不必要的逻辑
2026-07-07 09:55:00 +08:00
58 changed files with 3799 additions and 1370 deletions

View File

@@ -1,14 +1,29 @@
[package]
name = "ddddocr-rs"
version = "0.1.0"
[workspace]
resolver = "2"
members = [
"ddddocr-core", "ddddocr-ort",
"ddddocr-tract",
]
[workspace.package]
version = "0.2.1"
edition = "2024"
license = "MIT OR Apache-2.0"
[dependencies]
tract-onnx = { version = "0.21.10" }
anyhow = "1.0.102"
[workspace.dependencies]
tract-onnx = "0.23.4"
tract-linalg = { version = "0.23.4",features = ["multithread-mm"]}
ort = "2.0.0-rc.12"
ndarray = "0.17.2"
image = "0.25.10"
base64 = "0.22.1"
imageproc = { version = "0.26.2", default-features = true }
serde = { version = "1.0.228", features = ["derive"] }
serde_json = "1.0.150"
anyhow = "1.0.102"
thiserror = "1.0" # 刚好可以开始接入你需要的标准库错误处理
tracing = "0.1.44" # 埋入日志打点(后续需要继续优化,现在只是简单尝试)

16
ddddocr-core/Cargo.toml Normal file
View File

@@ -0,0 +1,16 @@
[package]
name = "ddddocr-core"
version = { workspace = true }
edition = { workspace = true }
license = { workspace = true }
[dependencies]
ndarray = { workspace = true } # 继承自工作空间
base64 = { workspace = true }
image = { workspace = true }
imageproc = { workspace = true }
thiserror = { workspace = true } # 刚好可以开始接入你需要的标准库错误处理
tracing={workspace = true}

6
ddddocr-core/src/det.rs Normal file
View File

@@ -0,0 +1,6 @@
mod builder;
mod executor;
pub use builder::DetBuilder;
pub use executor::{DetectionResult, Detector};
// pub use ddddocr_tract::det::session::DetSession;

View File

@@ -0,0 +1,11 @@
use crate::det::executor::Detector;
// use ddddocr_tract::det::session::DetSession;
use crate::traits::DetEngine;
#[derive(Default)]
pub struct DetBuilder;
impl DetBuilder {
fn build<E: DetEngine>(self, session: &E) -> Detector<'_> {
Detector { session }
}
}

View File

@@ -1,11 +1,12 @@
use crate::models::loader::{ModelLoader, ModelSession, ModelType};
use anyhow::{Context, Result};
use crate::error::{Result, TensorError};
use image::{DynamicImage, GenericImageView, imageops::FilterType};
use tract_onnx::prelude::tract_ndarray::{Array2, Array3, Array4, Axis, prelude::*, s};
use tract_onnx::prelude::{Graph, RunnableModel, Tensor, TypedFact, TypedOp, tvec};
use ndarray::{Array2, Array3, Array4, Axis, prelude::*, s};
use std::fmt;
// use tract_onnx::prelude::{Tensor};
// use ddddocr_tract::det::session::DetSession;
use crate::{DetBuilder, DetOutput, OcrBuilder};
use crate::traits::DetEngine;
#[derive(Debug, Clone, Copy)]
pub struct DetectionResult {
pub x1: i32,
@@ -16,30 +17,36 @@ pub struct DetectionResult {
pub class_id: u32,
}
pub struct Det {
session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
}
impl ModelSession for Det {
fn get_model_type(&self) -> ModelType {
todo!()
}
fn desc(&self) -> String {
"Detection Model 加载成功".to_string()
impl fmt::Display for DetectionResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
// 结构体只管自己这一行怎么显示,不用管外部的索引 [i]
write!(
f,
"x1={}, y1={}, x2={}, y2={}, 分数={:.4}, 类别ID={}",
self.x1, self.y1, self.x2, self.y2, self.score, self.class_id
)
}
}
impl Det {
pub fn new(model_path: String) -> Result<Self, anyhow::Error> {
let session = ModelLoader::load_model(&model_path)?.session;
Ok(Self { session })
pub struct Detector<'a> {
pub(crate) session: &'a dyn DetEngine,
}
impl<'a> Detector<'a> {
pub fn new(session: &'a dyn DetEngine) -> Self {
Detector { session }
}
pub fn builder() -> DetBuilder {
DetBuilder::default()
}
}
impl<'a> Detector<'a> {
pub fn predict(&self, image: &DynamicImage) -> Result<Vec<DetectionResult>> {
// Rust 中通常在调用层处理文件/PIL转换这里直接进入核心逻辑
self.get_bbox(image)
Ok(self.get_bbox(image)?)
}
/// 2. preproc: 纯 Rust 实现 (替代 OpenCV)
fn preproc(&self, image: &DynamicImage, input_size: (u32, u32)) -> Result<(Tensor, f32)> {
fn preproc(&self, image: &DynamicImage, input_size: (u32, u32)) -> (Array4<f32>, f32) {
let (target_h, target_w) = input_size;
let (img_w, img_h) = image.dimensions();
@@ -73,12 +80,11 @@ impl Det {
// BGR 赋值
array[[0, 0, y, x]] = slice[idx + 2] as f32; // B
array[[0, 1, y, x]] = slice[idx + 1] as f32; // G
array[[0, 2, y, x]] = slice[idx] as f32; // R
array[[0, 2, y, x]] = slice[idx] as f32; // R
}
}
Ok((array.into(), r))
(array, r)
}
/// 3. demo_postprocess (逻辑与 Python 一致)
@@ -236,19 +242,22 @@ impl Det {
.collect()
}
/// 6. get_bbox (完全解耦 OpenCV)
pub fn get_bbox(&self, dynamic_img: &DynamicImage) -> Result<Vec<DetectionResult>> {
pub fn get_bbox(
&self,
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))?;
let (input_tensor, ratio) = self.preproc(dynamic_img, (416, 416));
// tract 推理
let outputs = self.session.run(tvec!(input_tensor.into()))?;
let output_array = outputs[0]
.to_array_view::<f32>()?
.to_owned()
.into_dimensionality::<Ix3>()?;
// let outputs = self.session.session.run(tvec!(input_tensor.into()))?;
let outputs = self.session.inference(input_tensor)?;
// let output_array = outputs[0]
// 2. 无缝、安全地解包出标准 3维 矩阵
let DetOutput::Detection(output_array) = outputs;
let predictions = self.demo_postprocess(output_array, (416, 416));
let pred = predictions.slice(s![0, .., ..]);
@@ -256,9 +265,14 @@ impl Det {
let boxes = pred.slice(s![.., 0..4]);
let obj_conf = pred.slice(s![.., 4..5]);
let cls_conf = pred.slice(s![.., 5..]);
let obj_broadcast = obj_conf
.broadcast(cls_conf.dim())
.context("ndarray broadcasting failed for scores calculation")?;
let obj_broadcast =
obj_conf
.broadcast(cls_conf.dim())
.ok_or_else(|| TensorError::DimensionMismatch {
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..]);
@@ -273,17 +287,15 @@ impl Det {
let detections = self.multiclass_nms(&boxes_xyxy, &scores, 0.45, 0.1);
let final_results = detections
.into_iter()
.map(|d| {
DetectionResult{
x1: (d[0] as i32).max(0).min(orig_w as i32),
y1: (d[1] as i32).max(0).min(orig_h as i32),
x2: (d[2] as i32).max(0).min(orig_w as i32),
y2: (d[3] as i32).max(0).min(orig_h as i32),
score: d[4],
class_id: d[5] as u32,
}
.map(|d| DetectionResult {
x1: (d[0] as i32).max(0).min(orig_w as i32),
y1: (d[1] as i32).max(0).min(orig_h as i32),
x2: (d[2] as i32).max(0).min(orig_w as i32),
y2: (d[3] as i32).max(0).min(orig_h as i32),
score: d[4],
class_id: d[5] as u32,
})
.collect();
Ok(final_results )
Ok(final_results)
}
}

274
ddddocr-core/src/error.rs Normal file
View File

@@ -0,0 +1,274 @@
pub(crate) const MODEL_DOWNLOAD_HELP: &str = "\
================================================================================
[ddddocr-rust] 错误:未找到默认的模型文件!
--------------------------------------------------------------------------------
由于打包体积限制,本库未内置 ONNX 模型。请按照以下步骤操作:
1. 前往官方 GitHub 下载对应的模型权重:
- OCR 模型: https://github.com/sml2h3/ddddocr/raw/master/ddddocr/common_sml2h3_f32.onnx
- DET 模型: https://github.com/sml2h3/ddddocr/raw/master/ddddocr/common_det.onnx
2. 配置加载方式(二选一):
A. 【推荐】设置环境变量指向您下载的文件:
Linux/macOS: export DDDD_OCR_MODEL=\"/path/to/common_sml2h3_f32.onnx\"
Windows (CMD): set DDDD_OCR_MODEL=C:\\path\\to\\common_sml2h3_f32.onnx
Windows (PowerShell): $env:DDDD_OCR_MODEL=\"C:\\path\\to\\common_sml2h3_f32.onnx\"
B. 或者直接将模型文件重命名并放置在您运行程序的“当前工作目录”或“可执行文件同级目录”下。
================================================================================";
use thiserror::Error;
//
// #[derive(Error, Debug)]
// pub enum DdddError {
// // 【新增】专门处理文件读取、路径不存在等原生 I/O 错误
// #[error("系统网络或文件 I/O 异常: {0}")]
// Io(#[from] std::io::Error),
//
// #[error("图像预处理失败: {0}")]
// PreprocessError(#[from] ImagePreprocessReason),
//
// #[error("模型推理引擎内部发生异常: {0}")]
// EngineError(#[from] anyhow::Error),
//
// #[error("CTC 解码错误: {0}")]
// DecodeError(String),
//
// #[error("维度转换失败,预期维度 {expected},实际形状为 {actual:?}")]
// DimensionMismatch {
// expected: String,
// actual: Vec<usize>,
// },
//
// #[error("内存不连续,无法执行零拷贝操作")]
// NonContiguousMemory,
//
// #[error("未知的模型输出格式")]
// UnknownOutputFormat,
//
// #[error("解析节点 Fact 失败")]
// InternalError(String),
// }
//
// /// 专门服务于预处理的子错误枚举,保留全部底层上下文
// #[derive(Error, Debug)]
// pub enum ImagePreprocessReason {
// #[error("图片加载或文件 I/O 失败: {0}")]
// ImageIo(#[from] image::ImageError),
//
// #[error("图片转矩阵矩阵(ndarray)失败: {0}")]
// NdarrayError(#[from] ndarray::ShapeError),
//
// #[error("Base64 解码失败: {0}")]
// Base64(#[from] base64::DecodeError),
//
// #[error("Base64 头部格式不正确,缺少 ';base64,' 分隔符")]
// InvalidBase64Header,
//
// #[error("不支持的通道数: {0}")]
// UnsupportedChannels(usize),
//
// #[error("其他预处理错误: {0}")]
// Custom(String),
// }
/// 统一用我们自己的 DdddError 包装 Result
// pub type Result<T> = std::result::Result<T, DdddError>;
pub type Result<T, E = DdddError> = std::result::Result<T, E>;
// =====================================================================
// 1. 顶层全局 Error 分流器 (去 anyhow 化,完全基于标准库/自定义类型)
// =====================================================================
#[derive(Error, Debug)]
pub enum DdddError {
// /// 系统文件、网络等原生 I/O 异常 (高优先级自动转换)
// #[error("系统网络或文件 I/O 异常: {0}")]
// Io(#[from] std::io::Error),
/// 图像预处理阶段发生异常
#[error("图像预处理失败: {0}")]
Preprocess(#[from] ImagePreprocessError),
/// 推理引擎与张量操作阶段发生异常
#[error("推理与模型输入/输出张量异常: {0}")]
Inference(#[from] TensorError),
/// 算法后处理解码阶段发生异常
#[error("后处理解码错误: {0}")]
Decode(#[from] DecodeError),
/// 框架内部不可恢复的逻辑断言错误(例如解析节点 Fact 失败)
#[error("内部严重逻辑错误: {0}")]
Internal(String),
/// 【流派核心】接替 anyhow::Error 的用户自定义扩展错误
/// 承载任何第三方扩展、解密、特定预处理插件在执行时产生的自定义错误
#[error("用户自定义扩展错误: {0}")]
Other(#[source] Box<dyn std::error::Error + Send + Sync>),
}
// =====================================================================
// 2. 子领域 A: 图像预处理错误类型
// =====================================================================
#[derive(Error, Debug)]
pub enum ImagePreprocessError {
// #[error("图片加载或解码失败: {0}")]
// ImageIo(#[from] image::ImageError),
// image_io
#[error("图片转矩阵(ndarray)基础操作失败: {0}")]
Ndarray(#[from] ndarray::ShapeError),
// image_io
#[error("图像矩阵维度不合规!预期: {expected},实际图像形状: {actual:?}")]
InvalidDimensions {
expected: String,
actual: Vec<usize>,
},
// image_io
/// 从 ndarray 原始数据构建图像缓冲区时,缓冲区长度与分辨率/通道数不匹配
#[error(
"图像缓冲区长度不匹配!预期大小: {expected},实际大小: {actual} (分辨率: {width}x{height}, 通道数: {channels})"
)]
BufferLengthMismatch {
expected: usize,
actual: usize,
width: u32,
height: u32,
channels: usize,
},
// image_io
#[error("不支持的图像通道数: {0} (仅支持单通道灰度L、3通道RGB、4通道RGBA)")]
UnsupportedChannels(usize),
// ================= 新增:针对 HSV 和 Preset 的强类型错误 =================
/// HSV 颜色区间非法 (例如 H > 180 或 lower > upper)
#[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,
},
// #[error("Base64 解码失败: {0}")]
// Base64(#[from] base64::DecodeError),
//
// #[error("Base64 头部格式不正确,缺少 ';base64,' 分隔符")]
// InvalidBase64Header,
// #[error("其他预处理错误: {0}")]
// Other(String),
}
// =====================================================================
// 3. 子领域 B: 推理与张量操作错误类型
// =====================================================================
#[derive(Error, Debug)]
pub enum TensorError {
/// 替换原有的 anyhow::Error明确将 Tract/ONNX 引擎底层报错序列化为干净的 String
#[error("推理引擎内部发生异常: {0}")]
Engine(String),
/// 模型张量维度不匹配 (原有的顶层 DimensionMismatch 被优雅地归入本模块)
#[error("模型张量维度不匹配!预期: {expected},实际 Tensor 形状: {actual:?}")]
DimensionMismatch {
expected: String,
actual: Vec<usize>,
},
/// 新增:针对后处理 Logits 矩阵变形Reshape失败的精细化错误
/// 直接包装 ndarray::ShapeError保留强类型完美支持 match
#[error("OCR Logits 矩阵变形失败: {0}")]
LogitsDimensionMismatch(#[from] ndarray::ShapeError),
/// 张量内存布局不是连续的
#[error("内存不连续,无法执行零拷贝操作")]
NonContiguousMemory,
/// 模型的输出数据类型或格式不受支持
#[error("未知的模型输出格式")]
UnknownOutputFormat,
}
// =====================================================================
// 4. 子领域 C: 算法解码错误类型
// =====================================================================
#[derive(Error, Debug)]
pub enum DecodeError {
/// CTC 解码器解码过程中的逻辑报错
#[error("CTC 解码异常: {0}")]
Ctc(String),
}
// =====================================================================
// 5. 【自定义错误安全注入】不使用全局 `#[from]`,采用显式包装避免特化冲突
// =====================================================================
impl DdddError {
/// 提供类似 std::io::Error::new 的构造函数,方便手动且无痛地包装任意第三方错误
pub fn new<E>(error: E) -> Self
where
E: Into<Box<dyn std::error::Error + Send + Sync>>,
{
DdddError::Other(error.into())
}
// -----------------------------------------------------------------
// 2.3 优化提供一键判断与转换的快捷方法Downcasting Helpers
// -----------------------------------------------------------------
/// 快速判断是否是系统 I/O 错误
// pub fn is_io_error(&self) -> bool {
// matches!(self, DdddError::Io(_))
// }
/// 尝试将错误转换为引用形式的 `std::io::Error`
// pub fn as_io_error(&self) -> Option<&std::io::Error> {
// match self {
// DdddError::Io(err) => Some(err),
// _ => None,
// }
// }
/// 快速判断是否是预处理阶段的图片维度不合规错误
pub fn is_invalid_dimensions(&self) -> bool {
matches!(
self,
DdddError::Preprocess(ImagePreprocessError::InvalidDimensions { .. })
)
}
/// 快速判断是否是因为图片通道数不合规导致的失败
pub fn is_unsupported_channels(&self) -> bool {
matches!(
self,
DdddError::Preprocess(ImagePreprocessError::UnsupportedChannels(_))
)
}
// 提取出底层最原始的那个错误(无论是 IO、预处理、推理、还是第三方扩展错误
// 方便外层统一打印更深层的 `source` 链条
// pub fn source_error(&self) -> Option<&(dyn std::error::Error + 'static)> {
// use std::error::Error;
// match self {
// // DdddError::Io(err) => Some(err),
// DdddError::Preprocess(err) => Some(err),
// DdddError::Inference(err) => Some(err),
// DdddError::Decode(err) => Some(err),
// DdddError::Other(err) => Some(err.as_ref()),
// DdddError::Internal(_) => None, // Internal 内部目前只有 String没有底层的 Error source
// }
// }
}

24
ddddocr-core/src/lib.rs Normal file
View File

@@ -0,0 +1,24 @@
mod det;
pub mod error;
mod ocr;
mod slide;
pub mod utils;
pub mod types;
pub mod traits;
pub use crate::det::{DetBuilder, DetectionResult, Detector};
pub use crate::ocr::{Charset, ModelMetadata, Normalization, Ocr, OcrBuilder, OcrResult, Resize};
pub use crate::slide::{SlideResult, Slider};
// DetSession
pub enum OcrOutput {
Indices(ndarray::Array1<i64>), // 拥有完整所有权的 1维数组可任意传递和返回
Logits(ndarray::Array2<f32>),
}
/// 2. 目标检测专属的、编译期安全的输出枚举
pub enum DetOutput {
Detection(ndarray::Array3<f32>), // 拥有完整所有权的 2维矩阵可任意传递和返回
}

13
ddddocr-core/src/ocr.rs Normal file
View File

@@ -0,0 +1,13 @@
mod builder;
mod charset;
mod color_filter;
mod executor;
mod metadata;
mod token_filter;
pub use builder::OcrBuilder;
pub use charset::Charset;
pub use executor::{Ocr, OcrResult};
pub use metadata::{ModelMetadata, Normalization, Resize};
pub use token_filter::TokenFilter;
// pub use ddddocr_tract::session::OcrSession;

View File

@@ -0,0 +1,74 @@
use crate::ocr::executor::Ocr;
// use ddddocr_tract::session::OcrSession;
use crate::traits::OcrEngine;
use crate::ocr::color_filter::ColorFilter;
use crate::ocr::token_filter::TokenFilter;
#[derive(Default)]
pub struct OcrBuilder {
/// 是否修复PNG格式问题
png_fix: bool,
/// 是否返回概率信息
probability: bool,
/// 颜色过滤:保留的颜色列表
color_filter: Option<Box<dyn ColorFilter + Send + Sync>>,
/// 字符集范围
charset_restrict: Option<Box<dyn TokenFilter + Send + Sync>>,
}
impl OcrBuilder {
// 初始化任务,设置默认参数
pub fn new() -> Self {
Self {
png_fix: false, // 默认值
probability: false,
color_filter: None,
charset_restrict: None,
}
}
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,
{
self.color_filter = Some(Box::new(filter));
self
}
pub fn charset_restrict<T>(mut self, restrict: T) -> Self
where
T: TokenFilter + Send + Sync + 'static,
{
self.charset_restrict = Some(Box::new(restrict));
self
}
pub fn runner<E: OcrEngine>(self, runtime: &E) -> Ocr<'_> {
// 1. 原地解析颜色过滤器
let final_color_ranges = match &self.color_filter {
Some(filter) => filter.collect_to_vec(),
None => Ok(None),
};
// 2. 原地解析字符集过滤
let tokens = &runtime.metadata().charset.tokens;
let final_charset_indices = match &self.charset_restrict {
Some(restrict) => restrict.apply_to_charset(tokens),
None => None,
};
// Ocr::new(session, self)
Ocr {
runtime,
png_fix: self.png_fix, // 原地解构出来
probability: self.probability,
final_color_ranges,
final_charset_indices,
}
}
}

View File

@@ -0,0 +1,66 @@
use std::borrow::Cow;
use std::collections::HashMap;
// ==========================================
// 3. 字符集核心结构体 (重命名为 Charset)
// ==========================================
#[derive(Debug, Clone)]
pub struct Charset {
// 使用 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 {
// 内部底层统一收拢构造
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 {
tokens,
char_to_idx,
}
}
// --- 业务策略方法 ---
/// 将字符转为索引,不存在返回 -1 (保持与原 Python 库行为一致)
pub fn char_to_index(&self, char_str: &str) -> i32 {
if let Some(&idx) = self.char_to_idx.get(char_str) {
idx as i32
} else {
-1
}
}
/// 将索引转为字符引用,零拷贝。若越界返回 None
pub fn index_to_char_ref(&self, index: usize) -> Option<&str> {
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()
}
pub fn size(&self) -> usize {
self.tokens.len()
}
}
// ==========================================
// 4. 标准 Display 接口实现 (对应 __str__)
// ==========================================
impl std::fmt::Display for Charset {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "Charset [Total Size: {}", self.size(),)
}
}

View File

@@ -1,30 +1,30 @@
use std::str::FromStr;
use anyhow::anyhow;
use crate::error::{ImagePreprocessError, Result};
use crate::utils::image_processor::rgb_to_opencv_hsv;
use image::{DynamicImage, ImageBuffer, Rgb};
use crate::utils::cv_ops::rgb_to_opencv_hsv;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub struct HsvRange {
pub lower: (u8, u8, u8), // (H, S, V)
pub upper: (u8, u8, u8), // (H, S, V)
}
use std::str::FromStr;
/// 核心区间判定辅助函数
#[inline(always)]
fn is_pixel_matched(ranges: &[HsvRange], h: u8, s: u8, v: u8) -> bool {
ranges.iter().any(|range| {
h >= range.lower.0 && h <= range.upper.0 &&
s >= range.lower.1 && s <= range.upper.1 &&
v >= range.lower.2 && v <= range.upper.2
h >= range.lower.0
&& h <= range.upper.0
&& s >= range.lower.1
&& s <= range.upper.1
&& v >= range.lower.2
&& v <= range.upper.2
})
}
pub fn filter_image(image: &DynamicImage, hsv_ranges: &[HsvRange]) -> anyhow::Result<DynamicImage> {
pub fn apply_to_image(
image: &DynamicImage,
hsv_ranges: &[HsvRange],
) -> Result<DynamicImage, ImagePreprocessError> {
// 1. 统一转换为连续内存的 RGB8 缓冲区 (对应 Python 的 Image 到 RGB/BGR 数组转换)
let rgb_img = image.to_rgb8();
let (width, height) = rgb_img.dimensions();
let mut raw_pixels = rgb_img.into_raw();
let actual_len = raw_pixels.len();
let expected_len = (width as usize) * (height as usize) * 3;
// 2. 密集计算核心:原地流式迭代修改
// 每次取出 3 个 u8 字节,分别代表 [R, G, B],无多余掩膜矩阵内存分配
for chunk in raw_pixels.chunks_exact_mut(3) {
@@ -46,10 +46,23 @@ pub fn filter_image(image: &DynamicImage, hsv_ranges: &[HsvRange]) -> anyhow::Re
// 3. 将扁平字节数组重新打包回 DynamicImage 容器
let filtered_buffer = ImageBuffer::<Rgb<u8>, Vec<u8>>::from_raw(width, height, raw_pixels)
.ok_or_else(|| anyhow!("图像缓冲重新组装失败,维度与数据大小不匹配"))?;
// .ok_or_else(|| anyhow!("图像缓冲重新组装失败,维度与数据大小不匹配"))?;
.ok_or_else(|| ImagePreprocessError::BufferLengthMismatch {
expected: expected_len,
actual: actual_len,
width,
height,
channels: 3,
})?;
Ok(DynamicImage::ImageRgb8(filtered_buffer))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub struct HsvRange {
pub lower: (u8, u8, u8), // (H, S, V)
pub upper: (u8, u8, u8), // (H, S, V)
}
impl HsvRange {
pub const fn new(lower: (u8, u8, u8), upper: (u8, u8, u8)) -> Self {
Self { lower, upper }
@@ -58,15 +71,22 @@ impl HsvRange {
impl HsvRange {
/// 验证当前 HSV 范围是否合法
/// 对应 Python 逻辑H 在 0-180S/V 在 0-255且下界 <= 上界
pub fn validate(&self) -> Result<(), String> {
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("H通道值必须在 0-180 范围内".to_string());
return Err(ImagePreprocessError::InvalidHsvRange(
"H通道值必须在 0-180 范围内".to_string(),
));
}
// 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());
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(),
));
}
Ok(())
@@ -87,30 +107,63 @@ pub enum ColorPreset {
Custom(Vec<HsvRange>),
}
impl ColorPreset {
impl ColorPreset {
/// 纯裸数据定义,没有任何结构体包装,干净利落
/// 返回值:(范围数量, 范围数组)
/// 完美的零成本抽象:利用常量提升将数据直接打入只读数据段 (.rodata)
pub fn matches(&self) -> &[HsvRange] {
match self {
ColorPreset::Red => &[
HsvRange { lower: (0, 50, 50), upper: (10, 255, 255) },
HsvRange { lower: (170, 50, 50), upper: (180, 255, 255) },
HsvRange {
lower: (0, 50, 50),
upper: (10, 255, 255),
},
HsvRange {
lower: (170, 50, 50),
upper: (180, 255, 255),
},
],
ColorPreset::Blue => &[HsvRange { lower: (100, 50, 50), upper: (130, 255, 255) }],
ColorPreset::Green => &[HsvRange { lower: (40, 50, 50), upper: (80, 255, 255) }],
ColorPreset::Yellow => &[HsvRange { lower: (20, 50, 50), upper: (40, 255, 255) }],
ColorPreset::Orange => &[HsvRange { lower: (10, 50, 50), upper: (20, 255, 255) }],
ColorPreset::Purple => &[HsvRange { lower: (130, 50, 50), upper: (170, 255, 255) }],
ColorPreset::Cyan => &[HsvRange { lower: (80, 50, 50), upper: (100, 255, 255) }],
ColorPreset::Black => &[HsvRange { lower: (0, 0, 0), upper: (180, 255, 50) }],
ColorPreset::White => &[HsvRange { lower: (0, 0, 200), upper: (180, 30, 255) }],
ColorPreset::Gray => &[HsvRange { lower: (0, 0, 50), upper: (180, 30, 200) }],
ColorPreset::Blue => &[HsvRange {
lower: (100, 50, 50),
upper: (130, 255, 255),
}],
ColorPreset::Green => &[HsvRange {
lower: (40, 50, 50),
upper: (80, 255, 255),
}],
ColorPreset::Yellow => &[HsvRange {
lower: (20, 50, 50),
upper: (40, 255, 255),
}],
ColorPreset::Orange => &[HsvRange {
lower: (10, 50, 50),
upper: (20, 255, 255),
}],
ColorPreset::Purple => &[HsvRange {
lower: (130, 50, 50),
upper: (170, 255, 255),
}],
ColorPreset::Cyan => &[HsvRange {
lower: (80, 50, 50),
upper: (100, 255, 255),
}],
ColorPreset::Black => &[HsvRange {
lower: (0, 0, 0),
upper: (180, 255, 50),
}],
ColorPreset::White => &[HsvRange {
lower: (0, 0, 200),
upper: (180, 30, 255),
}],
ColorPreset::Gray => &[HsvRange {
lower: (0, 0, 50),
upper: (180, 30, 200),
}],
ColorPreset::Custom(ranges) => ranges,
}
}
/// 校验逻辑:在这里实现完美的“责任分离”
pub fn validate(&self) -> Result<(), String> {
pub fn validate(&self) -> Result<(), ImagePreprocessError> {
match self {
// 1. 快捷变体完全绕过根本不校验0 运行时开销放行!
ColorPreset::Custom(ranges) => {
@@ -126,7 +179,7 @@ impl ColorPreset {
}
impl FromStr for ColorPreset {
type Err = String;
type Err = ImagePreprocessError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"red" => Ok(ColorPreset::Red),
@@ -139,7 +192,8 @@ impl FromStr for ColorPreset {
"black" => Ok(ColorPreset::Black),
"white" => Ok(ColorPreset::White),
"gray" => Ok(ColorPreset::Gray),
_ => Err(format!("不支持的颜色预设: {}", s)),
// _ => Err(format!("不支持的颜色预设: {}", s)),
_ => Err(ImagePreprocessError::UnknownColorPreset(s.to_string())),
}
}
}
@@ -160,12 +214,12 @@ pub trait ColorFilter {
fn estimated_count(&self) -> usize;
/// 将自身的有效约束平铺追加到统一目标容器中
/// 验证当前过滤器是否合法默认直接放行Ok(())
fn validate_self(&self) -> Result<(), String> {
fn validate_self(&self) -> Result<(), ImagePreprocessError> {
Ok(())
}
/// 【新扩展的架构方法】将自身安全的合并到已有的普通容器中,并完成去重和排序
/// 完美的责任分离Builder 不再需要关心怎么分配内存、怎么排序去重
fn collect_to_vec(&self) -> Result<Option<Vec<HsvRange>>, String> {
fn collect_to_vec(&self) -> Result<Option<Vec<HsvRange>>, ImagePreprocessError> {
// 1. 触发自检
self.validate_self()?;
@@ -197,14 +251,13 @@ impl ColorFilter for ColorPreset {
// 直接获取切片长度
self.matches().len()
}
fn validate_self(&self) -> Result<(), String> {
fn validate_self(&self) -> Result<(), ImagePreprocessError> {
// 直接调用我们在第一步中为 ColorPreset 实现的精细化分流校验
// 快捷变体在这里会直接返回 Ok(()), 只有 Custom 才会去真正校验
self.validate()
}
}
/// 多路颜色“或”逻辑组合子(并集网络)
pub struct MultiOrColorRestrict<'a> {
pub filters: Vec<&'a dyn ColorFilter>,
@@ -222,7 +275,7 @@ impl<'a> ColorFilter for MultiOrColorRestrict<'a> {
self.filters.iter().map(|f| f.estimated_count()).sum()
}
fn validate_self(&self) -> Result<(), String> {
fn validate_self(&self) -> Result<(), ImagePreprocessError> {
// 递归政审:只要其中一个子过滤器校验失败(比如某个 Custom 变体非法),立刻熔断
for f in &self.filters {
f.validate_self()?;
@@ -246,4 +299,3 @@ macro_rules! color_any_of {
}
};
}

View File

@@ -1,26 +1,25 @@
use crate::charset::{TokenFilter, ValidationCtx};
use crate::model_metadata::{ModelMetadata, Resize};
use crate::models::base::ModelArgs;
use crate::models::loader::{ModelLoader, ModelSession, ModelType};
use crate::utils::color_filter::{ColorFilter, HsvRange, filter_image};
use crate::utils::image_io::png_rgba_white_preprocess;
use crate::utils::image_processor::{convert_to_grayscale, resize_image};
use anyhow::Context;
use anyhow::{Result, anyhow};
use image::{DynamicImage, ImageBuffer, Rgb};
use serde::Serialize;
use std::borrow::Cow;
use std::collections::HashSet;
use std::fmt;
use tract_onnx::prelude::tract_ndarray::{ArrayView2, Ix2, s};
use tract_onnx::prelude::{
DatumType, Graph, IntoTensor, RunnableModel, Tensor, TypedFact, TypedOp, tract_ndarray, tvec,
};
// 引入 cv_ops 模块中的 OpenCV HSV 转换算子
use crate::utils::cv_ops::rgb_to_opencv_hsv;
use crate::ocr::metadata::Resize;
/// 推理最终输出的强类型外壳(完全 Owned无任何生命周期可直接转 JSON
#[derive(Debug, Clone, Serialize)]
use crate::ocr::color_filter::{HsvRange, apply_to_image};
// use ddddocr_tract::session::{ModelOutput, OcrSession};
use crate::utils::image_convert::png_rgba_white_preprocess;
use crate::utils::image_processor::{convert_to_grayscale, resize_image};
use image::DynamicImage;
use std::borrow::Cow;
use std::fmt;
// use tract_onnx::prelude::tract_ndarray::{ Ix2, s};
// use tract_onnx::prelude::{DatumType, Tensor, tract_ndarray};
// !!!【核心纠正】:彻底弃用 tract_ndarray全线转用标准 ndarray
use ndarray::ArrayView2;
// pub enum ModelOutput {
// Indices(ndarray::Array1<i64>), // 拥有完整所有权的 1维数组可任意传递和返回
// Logits(ndarray::Array2<f32>), // 拥有完整所有权的 2维矩阵可任意传递和返回
// }
use crate::error::{ImagePreprocessError, Result, TensorError};
use crate::{OcrBuilder, OcrOutput};
use crate::traits::OcrEngine;
use tracing::{ warn};
#[derive(Debug, Clone)]
pub enum OcrResult {
/// 纯文本分支(对应 probability = false
Text(String),
@@ -104,108 +103,52 @@ impl fmt::Display for OcrResult {
}
}
pub struct Ocr {
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
pub model_metadata: ModelMetadata,
}
impl ModelSession for Ocr {
fn get_model_type(&self) -> ModelType {
todo!("使用thiserror作为错误处理的库,thiserror 专门用于开发库Library");
}
fn desc(&self) -> String {
"Ocr Model 加载成功".to_string()
}
}
impl Ocr {
pub fn new(model_path: String, model_metadata: ModelMetadata) -> Result<Self, anyhow::Error> {
let session = ModelLoader::load_model(&model_path)?.session;
Ok(Self {
session,
model_metadata,
})
}
/// 对应 Python 的 _inference
fn inference(&self, tensor: Tensor) -> anyhow::Result<Tensor> {
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
// let result = self.session.run(tvec!(tensor.into()))?;
let mut result = self
.session
.run(tvec!(tensor.into()))
.context("执行模型推理失败")?;
println!("模型输出原始数据: {:?}", result);
Ok(result.swap_remove(0).into_tensor())
}
pub fn predictor(&'_ self) -> OcrPredictor<'_> {
OcrPredictor::new(self)
}
}
pub struct OcrPredictor<'a> {
ocr: &'a Ocr,
/// 是否修复PNG格式问题
png_fix: bool,
/// 是否返回概率信息
probability: bool,
pub struct Ocr<'a> {
pub(crate) runtime: &'a dyn OcrEngine,
pub(crate) png_fix: bool,
pub(crate) probability: bool,
/// 颜色过滤:保留的颜色列表
color_filter: Result<Option<Vec<HsvRange>>, String>,
pub(crate) final_color_ranges: Result<Option<Vec<HsvRange>>, ImagePreprocessError>,
/// 字符集范围
charset_restrict: Option<Vec<usize>>,
pub(crate) final_charset_indices: Option<Vec<usize>>,
}
impl<'a> OcrPredictor<'a> {
impl<'a> Ocr<'a> {
// 初始化任务,设置默认参数
pub fn new(ocr: &'a Ocr) -> Self {
Self {
ocr,
pub fn new(runtime: &'a dyn OcrEngine) -> Self {
Ocr {
runtime,
png_fix: false, // 默认值
probability: false,
color_filter: Ok(None),
charset_restrict: None,
final_color_ranges: Ok(None),
final_charset_indices: None,
}
}
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 builder() -> OcrBuilder {
OcrBuilder::default()
}
pub fn color_filter(mut self, filter: &dyn ColorFilter) -> Self {
// 一句话把活全包了!错误信息无缝传递,完美熔断
match filter.collect_to_vec() {
Ok(new_ranges) => self.color_filter = Ok(new_ranges),
Err(err_msg) => self.color_filter = Err(err_msg), // 校验失败Builder 正式中毒
}
self
}
pub fn charset_restrict(mut self, restrict: &dyn TokenFilter) -> Self {
let charset = &self.ocr.model_metadata.charset;
let tokens = &charset.tokens;
self.charset_restrict = restrict.apply_to_charset(tokens);
self
}
}
impl<'a> OcrPredictor<'a> {
pub fn predict(self, image: &DynamicImage) -> anyhow::Result<OcrResult> {
println!("当前颜色过滤器状态: {:?}", self.color_filter);
impl<'a> Ocr<'a> {
pub fn predict(&self, image: &DynamicImage) -> Result<OcrResult> {
println!("当前颜色过滤器状态: {:?}", self.final_color_ranges);
// =====================================================================
// 管道节点 1: 颜色过滤流水线
// 使用 Cow (Copy-On-Write) 智能指针。
// 如果未开启过滤img_cow 内部只是持有原图的【只读借用】,发生【零内存分配】!
// =====================================================================
let img_cow = match &self.color_filter {
let img_cow = match &self.final_color_ranges {
Err(err_msg) => {
return Err(anyhow::anyhow!(
"颜色过滤器初始化失败,全链路短路: {}",
err_msg
));
// return Err(anyhow::anyhow!(
// "颜色过滤器初始化失败,全链路短路: {}",
// err_msg
// ));
return Err(ImagePreprocessError::FilterConfigInvalid(
err_msg.to_string(),
))?;
}
Ok(None) => {
// 核心优化点:直接借用原图,不发生任何克隆
@@ -213,34 +156,34 @@ impl<'a> OcrPredictor<'a> {
}
Ok(Some(ranges)) => {
// 只有真正需要过滤时,才在内部提取像素并生成清洗后的 Owned 新图
let filtered_img = filter_image(image, ranges)?;
let filtered_img = apply_to_image(image, ranges)?;
Cow::Owned(filtered_img)
}
};
let tensor = self.preprocess_image(&img_cow)?;
let raw_tensor = self.ocr.inference(tensor)?;
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 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)
}
/// 对应 Python 的 _preprocess_image
/// 负责:透明背景修复 -> 灰度化 -> 按比例 Resize -> 归一化 -> 4维张量转换
fn preprocess_image(&self, img: &DynamicImage) -> anyhow::Result<Tensor> {
fn preprocess_image(&self, img: &DynamicImage) -> Result<ndarray::Array4<f32>,ImagePreprocessError> {
// 1. 获取模型元数据配置
let meta = &self.ocr.model_metadata;
let meta = self.runtime.metadata();
let norm = &meta.normalization; // 获取归一化器
// A. 修复 PNG 透明背景 (内部逻辑你之前已实现)
@@ -270,12 +213,12 @@ impl<'a> OcrPredictor<'a> {
let resized_img = resize_image(&current_img, target_w, target_h);
// 4. 管道节点 3: 颜色通道转换(单通道灰度 vs 三通道 RGB与 4D 张量填充
let tensor = match meta.channel {
let array4 = match meta.channel {
// --- 情况 A: 单通道(灰度图),对应 Python 的 len(shape) == 2 展开 ---
1 => {
let gray_img = convert_to_grayscale(&resized_img);
let array = tract_ndarray::Array4::from_shape_fn(
let array = 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;
@@ -284,14 +227,14 @@ impl<'a> OcrPredictor<'a> {
norm.normalize(pixel)
},
);
Tensor::from(array)
array
}
// --- 情况 B: 三通道RGB对应 Python 的 transpose(2, 0, 1) 的 CHW 布局 ---
3 => {
let rgb_img = resized_img.to_rgb8();
let array = tract_ndarray::Array4::from_shape_fn(
let array = 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;
@@ -300,13 +243,19 @@ impl<'a> OcrPredictor<'a> {
norm.normalize(pixel)
},
);
Tensor::from(array)
// Tensor::from(array)
array
}
_ => return Err(anyhow::anyhow!("不支持的通道数配置: {}", meta.channel)),
// _ => return Err(anyhow::anyhow!("不支持的通道数配置: {}", meta.channel)),
_ => {
return Err(ImagePreprocessError::UnsupportedChannels(
meta.channel as usize,
));
}
};
Ok(tensor)
Ok(array4)
// Ok(tensor)
// let h = 64u32;
// let w = (current_img.width() as f32 * (h as f32 / current_img.height() as f32)) as u32;
@@ -327,14 +276,66 @@ impl<'a> OcrPredictor<'a> {
//
// Ok(tensor)
}
// 这段代码未来直接放入 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)?;
let final_text = self.ctc_decode_to_string(slice);
if self.probability {
Ok(OcrResult::Probability {
text: final_text,
probabilities: vec![],
confidence: 1.0,
})
} else {
Ok(OcrResult::Text(final_text))
}
}
OcrOutput::Logits(matrix_view) => {
// 对应你原来的 process_f32_tensor
// 注意:此时的 matrix_view 已经是干净的标准的 ndarray::Array2<f32>,且保证是 [Steps, Classes] 2D 形状
if self.probability {
let (probabilities_list, confidence, predicted_indices) =
self.compute_f32_full_probability(matrix_view.view());
let final_text = self.ctc_decode_to_string(&predicted_indices);
Ok(OcrResult::Probability {
text: final_text,
probabilities: probabilities_list,
confidence: confidence as f64,
})
} else {
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))
}
}
}
}
}
impl<'a> OcrPredictor<'a> {
impl<'a> Ocr<'a> {
fn is_valid_indices(&self, idx: usize) -> bool {
if idx >= self.ocr.model_metadata.charset.size() {
if idx >= self.runtime.metadata().charset.size() {
return false;
}
match &self.charset_restrict {
match &self.final_charset_indices {
Some(v) => v.binary_search(&idx).is_ok(),
None => true,
}
@@ -342,9 +343,9 @@ impl<'a> OcrPredictor<'a> {
/// 【按需延迟打印】:当用户真的需要“知道当前有哪些限制字符”时,一秒反查并打印
/// 这里的 &str 完美借用了自 tokens依然是彻底的零拷贝
pub fn valid_tokens(&self) -> Vec<&str> {
let charset = &self.ocr.model_metadata.charset;
let charset = &self.runtime.metadata().charset;
let tokens = &charset.tokens;
match &self.charset_restrict {
match &self.final_charset_indices {
Some(indices) => indices
.iter()
.filter_map(|&idx| tokens.get(idx).map(|cow| cow.as_ref()))
@@ -354,9 +355,9 @@ impl<'a> OcrPredictor<'a> {
}
}
pub fn valid_size(&self) -> usize {
match &self.charset_restrict {
match &self.final_charset_indices {
Some(indices) => indices.len(),
None => self.ocr.model_metadata.charset.tokens.len(),
None => self.runtime.metadata().charset.tokens.len(),
}
}
/// 变体 B 核心处理器:单次遍历 2D 视图,融合计算 Softmax、Argmax、置信度并输出概率大包
@@ -368,7 +369,7 @@ impl<'a> OcrPredictor<'a> {
let classes = matrix_view.ncols();
// 1. 预分配满额概率矩阵内存
let mut prob_matrix = tract_ndarray::Array2::<f32>::zeros((steps, classes));
let mut prob_matrix = ndarray::Array2::<f32>::zeros((steps, classes));
let mut predicted_indices = Vec::with_capacity(steps);
let mut confidence_sum = 0.0f32;
@@ -413,98 +414,98 @@ impl<'a> OcrPredictor<'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 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);
let charset = &self.ocr.model_metadata.charset;
let charset = &self.runtime.metadata().charset;
let tokens = &charset.tokens;
// let valid_indices = &charset.valid_indices;
@@ -532,7 +533,7 @@ impl<'a> OcrPredictor<'a> {
// 史诗级加速点:如果是 None说明没限制根本不进入分支直接放行
// 只有当有具体限制Some才去跑 4-5 次 CPU 寄存器级别的二分查找
if let Some(ref indices) = self.charset_restrict {
if let Some(ref indices) = self.final_charset_indices {
if indices.binary_search(&u_idx).is_err() {
continue;
}
@@ -541,8 +542,9 @@ impl<'a> OcrPredictor<'a> {
// 5. 字符映射
if let Some(char_str) = tokens.get(u_idx) {
res.push_str(char_str);
} else {
eprintln!("警告: 预测索引 {} 超出字符集范围", u_idx);
}
else {
warn!("警告: 预测索引 {} 超出字符集范围", u_idx);
}
}
res

View File

@@ -0,0 +1,86 @@
// =====================================================================
// 1. 辅助定义的枚举与结构体
// =====================================================================
use crate::ocr::Charset;
use std::borrow::Cow;
#[derive(Debug, Clone, Copy)]
pub enum Normalization {
/// 映射到 [0.0, 1.0] -> pixel / 255.0
ZeroToOne,
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
MinusOneToOne,
}
impl Normalization {
/// 统一归一化计算逻辑
#[inline(always)]
pub fn normalize(&self, pixel: f32) -> f32 {
match self {
Normalization::ZeroToOne => pixel / 255.0,
Normalization::MinusOneToOne => (pixel / 255.0 - 0.5) / 0.5,
}
}
}
/// 图像缩放策略枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Resize {
/// 固定宽高,例如 (64, 64)
Fixed(u32, u32),
/// 高度固定,宽度根据原始比例动态计算(对应 Python 的 [-1, H]
DynamicWidth(u32),
/// 单字识别的正方形切图(对应 Python 的 word 为 True 且 [-1, H]
Square(u32),
}
#[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,
resize: Resize,
channel: u8,
normalization: Normalization,
) -> Self {
Self {
charset,
word,
resize,
channel,
normalization,
}
}
// --- 优雅的工厂模式构造器 ---
/// 通用的静态切片转换构造器
pub fn from_static_slice(
slice: &[&'static str],
word: bool,
resize: Resize,
channel: u8,
normalization: Normalization,
) -> Self {
let tokens: Vec<Cow<'static, str>> = slice.iter().map(|&s| Cow::Borrowed(s)).collect();
Self {
charset: Charset::new(tokens),
word,
resize,
channel,
normalization,
}
}
}

View File

@@ -0,0 +1,146 @@
use std::borrow::Cow;
/// 字符集范围限制枚举
pub struct ValidationCtx<'a> {
pub text: &'a str, // 当前 Token 的文本内容
pub token_id: usize, // 当前 Token 的 ID 索引
}
/// 统一的约束接口
pub trait TokenFilter {
fn matches(&self, ctx: &ValidationCtx) -> bool;
/// 预估容量提示,帮助精准开辟 Vec 内存
fn estimated_capacity(&self) -> usize {
128
}
/// 【新引入的架构级核心方法】
/// 统一接管全量字符集的密集遍历、CTC Blank放行、去重、排序及空交集退化兜底
fn apply_to_charset(&self, tokens: &[Cow<str>]) -> Option<Vec<usize>> {
let mut has_any_match = false;
let estimated_capacity = self.estimated_capacity();
// 1. 精准开辟内存,完美利用容量提示,避免动态乱涨
let mut temp_indices = Vec::with_capacity(estimated_capacity.max(16));
// 2. 高性能原地单次流式迭代
for (idx, token) in tokens.iter().enumerate() {
let token_str = token.as_ref();
// 规则 A: CTC Blank 空字符串或 0 号索引无条件放行
if token_str.is_empty() || idx == 0 {
temp_indices.push(idx);
continue; // 关键:直接跳过,防止后续 matches 匹配成功导致重复 push 产生 Bug
}
// 规则 B: 组装无拷贝上下文
let ctx = ValidationCtx {
text: token_str,
token_id: idx,
};
// 规则 C: 路由到各自具体实现的特异性匹配中(如 Digit 判定、TopN 判定、组合子判定等)
if self.matches(&ctx) {
temp_indices.push(idx);
has_any_match = true;
}
}
// 3. 终极防御:如果整个模型字符集除了 Blank一个都没对上直接退化为 None全量识别
if !has_any_match {
println!("警告:当前限制策略与模型字符集完全没有交集!已自动恢复全量识别。");
None
} else {
// 4. 排序并去重,为 Ocr 引擎后续进行极其高频的『二分查找』筑起绝对安全的底层保障
temp_indices.sort_unstable();
temp_indices.dedup();
Some(temp_indices)
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CharRestrict {
Digit,
Lowercase,
Uppercase,
CustomList(Vec<String>),
}
impl TokenFilter for CharRestrict {
fn matches(&self, ctx: &ValidationCtx) -> bool {
match self {
Self::Digit => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_digit(),
Self::Lowercase => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_lowercase(),
Self::Uppercase => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_uppercase(),
Self::CustomList(vec) => vec.iter().any(|t| t == ctx.text),
}
}
fn estimated_capacity(&self) -> usize {
match self {
Self::Digit => 16,
Self::Lowercase | Self::Uppercase => 32,
Self::CustomList(vec) => vec.len() + 1,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IdRestrict {
TopN(usize),
IdRange(std::ops::Range<usize>),
IdList(Vec<usize>),
}
impl TokenFilter for IdRestrict {
fn matches(&self, ctx: &ValidationCtx) -> bool {
match self {
Self::TopN(n) => ctx.token_id < *n,
Self::IdRange(range) => range.contains(&ctx.token_id),
Self::IdList(vec) => vec.contains(&ctx.token_id),
}
}
fn estimated_capacity(&self) -> usize {
match self {
Self::TopN(n) => *n + 1,
// 2. IdRange标准标准库 Range 的长度
// 注意:因为范围可能是 1000..2000,它的 len() 返回的是 usize
Self::IdRange(range) => range.len() + 1,
// 3. IdListVec 里的元素个数
Self::IdList(vec) => vec.len() + 1,
}
}
}
/// 多路“或”逻辑组合子(支持 N 个规则无缝并集)
pub struct MultiOrRestrict<'a> {
pub filters: Vec<&'a dyn TokenFilter>,
}
impl<'a> TokenFilter for MultiOrRestrict<'a> {
fn matches(&self, ctx: &ValidationCtx) -> bool {
// 核心高阶函数:只要有一个过滤器命中,该 Token 即可放行
self.filters.iter().any(|f| f.matches(ctx))
}
fn estimated_capacity(&self) -> usize {
// 将所有过滤器的预估容量累加,作为最终容量参考
self.filters.iter().map(|f| f.estimated_capacity()).sum()
}
}
// =====================================================================
// 声明式宏:替代 `+` 运算符,解决组合扩展痛苦
// =====================================================================
#[macro_export]
macro_rules! any_of {
// 场景 A如果用户只传了一个规则免去构建 Vec 的开销,直接返回其引用
($only:expr) => {
&$only as &dyn $crate::TokenFilter
};
// 场景 B如果用户传入了多个规则自动织成一张静态组合网
($($filter:expr),+ $(,)?) => {
&$crate::MultiOrRestrict {
filters: vec![ $( &$filter as &dyn $crate::TokenFilter ),+ ]
}
};
}

View File

@@ -1,32 +1,40 @@
use crate::utils::cv_ops;
use crate::utils::cv_ops::{abs_diff, min_max_loc, ndarray_to_luma8, rgb_to_gray};
use crate::utils::image_io::image_to_ndarray;
use anyhow::{Context, Result, anyhow};
use image::{DynamicImage, GenericImageView};
use image::{ImageBuffer, Luma};
use crate::error::{ImagePreprocessError, Result};
use crate::utils::image_convert::{ColorMode, image_to_ndarray};
use crate::utils::image_processor;
use crate::utils::image_processor::{abs_diff, min_max_loc, ndarray_to_luma8, rgb_to_gray};
use image::DynamicImage;
use image::Luma;
use imageproc::contrast::{ThresholdType, threshold};
use imageproc::distance_transform::Norm;
use imageproc::edges::canny;
use imageproc::morphology::{close, open};
use imageproc::region_labelling::{Connectivity, connected_components};
use imageproc::template_matching::{MatchTemplateMethod, match_template};
use std::cmp::{max, min};
use tract_onnx::prelude::tract_ndarray::{Array2, Array3, ArrayView2, ArrayView3, Axis, s};
use ndarray::{ArrayView2, ArrayView3};
use std::fmt;
#[derive(Debug)]
pub struct SlideResult {
pub target: [i32; 2],
pub target_x: i32,
pub target_y: i32,
pub confidence: f64,
}
impl fmt::Display for SlideResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(f, "滑块匹配测试结果:")?;
writeln!(f, "检测坐标: [x: {}, y: {}]", self.target_x, self.target_y)?;
// 注意:这里保留 4 位小数,如果想让外部控制,也可以直接写 {:.4}
write!(f, "置信度: {:.4}", self.confidence)?;
Ok(())
}
}
pub struct Slide;
pub struct Slider;
impl Slide {
impl Slider {
pub fn new() -> Self {
Self
}
/// 对应 Python: slide_match 滑块匹配接口
pub fn slide_match(
&self,
@@ -34,10 +42,11 @@ impl Slide {
background_image: &DynamicImage,
simple_target: bool,
) -> Result<SlideResult> {
let target_array = image_to_ndarray(target_image);
let background_array = image_to_ndarray(background_image);
let target_array = image_to_ndarray(target_image, ColorMode::RGB)?;
let background_array = image_to_ndarray(background_image, ColorMode::RGB)?;
self.perform_slide_match(target_array.view(), background_array.view(), simple_target)
.map_err(Into::into)
}
/// 对应 Python: slide_comparison 差异比较接口
/// 用于比较带坑位的图片与原始背景图,定位差异点
@@ -47,52 +56,35 @@ impl Slide {
background_image: &DynamicImage,
) -> Result<SlideResult> {
// 1. 转换为 ndarray (HWC RGB)
let target_array = image_to_ndarray(target_image);
let background_array = image_to_ndarray(background_image);
let target_array = image_to_ndarray(target_image, ColorMode::RGB)?;
let background_array = image_to_ndarray(background_image, ColorMode::RGB)?;
// 2. 执行比较逻辑 (对应 _perform_slide_comparison)
self.perform_slide_comparison(target_array.view(), background_array.view())
.map_err(Into::into)
}
/// 对应 Python: _perform_slide_comparison
pub fn perform_slide_comparison(
&self,
target: ArrayView3<u8>,
background: ArrayView3<u8>,
) -> Result<SlideResult> {
// let (h, w, _) = target.dim();
// 1. 计算图像差异并灰度化 (对应 cv2.absdiff + cv2.cvtColor)
// 使用 OpenCV 标准权重公式0.299R + 0.587G + 0.114B
// let mut diff_buffer = ImageBuffer::new(w as u32, h as u32);
// for y in 0..h {
// for x in 0..w {
// let r_diff = (target[[y, x, 0]] as i16 - background[[y, x, 0]] as i16).abs() as f32;
// let g_diff = (target[[y, x, 1]] as i16 - background[[y, x, 1]] as i16).abs() as f32;
// let b_diff = (target[[y, x, 2]] as i16 - background[[y, x, 2]] as i16).abs() as f32;
//
// let gray_diff = (0.299 * r_diff + 0.587 * g_diff + 0.114 * b_diff) as u8;
// diff_buffer.put_pixel(x as u32, y as u32, Luma([gray_diff]));
// }
// }
) -> Result<SlideResult, ImagePreprocessError> {
// 1. 计算差异数组 (复用 cv2::absdiff)
let (th, tw, tc) = target.dim();
let (bh, bw, bc) = background.dim();
// 1. 比较模式下的严格尺寸校验
if th != bh || tw != bw || tc != bc {
return Err(anyhow!(
return Err(ImagePreprocessError::MismatchDimensions(format!(
"比较模式要求两张图分辨率与通道数完全一致Target: [{}x{}x{}], Background: [{}x{}x{}]",
tw,
th,
tc,
bw,
bh,
bc
));
tw, th, tc, bw, bh, bc
)));
}
if th == 0 || tw == 0 {
return Err(anyhow!("输入图像尺寸不能为0"));
return Err(ImagePreprocessError::InvalidDimensions {
expected: "输入图像尺寸不能为0".to_string(),
actual: vec![th, tw],
});
}
let diff_array = abs_diff(&target, &background);
@@ -118,11 +110,12 @@ impl Slide {
// // 统计每个标签出现的频率(即面积)
// 4. 寻找最大连通区域 (对应 findContours + max area)
if let Some(max_label) = cv_ops::find_contours_and_max(&labelled) {
if let Some(max_label) = image_processor::find_contours_and_max(&labelled) {
// 5. 计算最大区域的边界框 (对应 cv2.boundingRect)
let (x, y, w, h) = cv_ops::bounding_rect(&labelled, max_label);
let (x, y, w, h) = image_processor::bounding_rect(&labelled, max_label);
// 6. 计算中心点 (调用之前封装的 calculate_center)
let (center_x, center_y) = cv_ops::calculate_center((x, y), w as usize, h as usize);
let (center_x, center_y) =
image_processor::calculate_center((x, y), w as usize, h as usize);
Ok(SlideResult {
target: [center_x, center_y],
@@ -147,29 +140,31 @@ impl Slide {
target: ArrayView3<u8>,
background: ArrayView3<u8>,
simple_target: bool, // 增加这个参数
) -> Result<SlideResult> {
) -> Result<SlideResult, ImagePreprocessError> {
let (th, tw, tc) = target.dim();
let (bh, bw, bc) = background.dim();
// 1. 严格的鲁棒性校验(防止底层的 imageproc 算子崩溃)
if th == 0 || tw == 0 || bh == 0 || bw == 0 {
return Err(anyhow!("输入图像的宽度或高度不能为0"));
return Err(ImagePreprocessError::InvalidDimensions {
expected: "输入图像的宽度或高度不能为0".to_string(),
actual: vec![th, tw, tc],
});
}
if th > bh || tw > bw {
return Err(anyhow!(
"尺寸不匹配:滑块模板(target)尺寸 [{}x{}] 不能大于背景图(background) [{}x{}]",
tw,
th,
bw,
bh
));
return Err(ImagePreprocessError::TargetExceedsBackground {
// "尺寸不匹配:滑块模板(target)尺寸 [{}x{}] 不能大于背景图(background) [{}x{}]",
target_w: tw,
target_h: th,
bg_w: bw,
bg_h: bh,
});
}
if tc != bc {
return Err(anyhow!(
return Err(ImagePreprocessError::MismatchDimensions(format!(
"目标图与背景图的通道数不一致 (target: {}, bg: {})",
tc,
bc
));
tc, bc
)));
}
// 1. 统一灰度化
@@ -178,11 +173,11 @@ impl Slide {
if simple_target {
// 2a. 简单模式:直接在灰度图上匹配
self.simple_template_match(target_gray.view(), background_gray.view())
Ok(self.simple_template_match(target_gray.view(), background_gray.view()))
} else {
// 2b. 复杂模式:先提取边缘,再匹配
self.edge_based_match(target_gray.view(), background_gray.view())
Ok(self.edge_based_match(target_gray.view(), background_gray.view()))
}
}
/// 对应 Python: _simple_template_match
@@ -192,11 +187,8 @@ impl Slide {
&self,
target: ArrayView2<u8>,
background: ArrayView2<u8>,
) -> Result<SlideResult> {
) -> SlideResult {
// 1. 将 ndarray 转换为 imageproc 需要的 ImageBuffer (无拷贝或轻量转换)
// let (bh, bw) = background.dim();
// 转换逻辑 (假设你已经有方法转回 ImageBuffer)
let t_buf = ndarray_to_luma8(target);
let b_buf = ndarray_to_luma8(background);
@@ -216,16 +208,17 @@ impl Slide {
// 4. 计算中心点 (与 Python 逻辑完全一致)
let (th, tw) = target.dim();
let (center_x, center_y) = cv_ops::calculate_center(max_loc, tw as usize, th as usize);
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);
Ok(SlideResult {
SlideResult {
target: [center_x, center_y],
target_x: center_x,
target_y: center_y,
confidence: max_val as f64,
})
}
}
/// 对应 Python: _edge_based_match
@@ -234,7 +227,7 @@ impl Slide {
&self,
target: ArrayView2<u8>,
background: ArrayView2<u8>,
) -> Result<SlideResult> {
) -> SlideResult {
// 1. 将 ndarray 转换为 ImageBuffer
// 注意Canny 和 match_template 需要 ImageBuffer 格式
let t_buf = ndarray_to_luma8(target);
@@ -261,18 +254,19 @@ impl Slide {
// 5. 计算中心位置 (对齐 Python 逻辑)
// target_w, target_h 来自输入数组的维度
let (th, tw) = target.dim();
let (center_x, center_y) = cv_ops::calculate_center(max_loc, tw as usize, th as usize);
let (center_x, center_y) =
image_processor::calculate_center(max_loc, tw as usize, th as usize);
// 打印调试信息,方便与 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);
Ok(SlideResult {
SlideResult {
target: [center_x, center_y],
target_x: center_x,
target_y: center_y,
confidence: max_val as f64,
})
}
}
}

View File

@@ -0,0 +1,42 @@
use crate::error::TensorError;
use crate::types::{ModelInfo, TensorInfo};
use crate::{DetOutput, ModelMetadata, OcrOutput};
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>;
}
/// 核心层定义的统一推理引擎接口。
/// 未来的 ddddocr-tract 和 ddddocr-ort 都必须实现这个 Trait
pub trait InferenceEngine {
/// 关联类型:具体的 Session 需要声明自己到底产出什么枚举
type Output;
fn inference(
&self,
input_array: ndarray::Array4<f32>,
) -> crate::error::Result<Self::Output, TensorError>;
}
pub trait OcrEngine: InferenceEngine<Output = OcrOutput> + Info {
fn metadata(&self) -> &ModelMetadata;
}
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],
) -> crate::error::Result<Self::Session, Self::Error>;
}

46
ddddocr-core/src/types.rs Normal file
View File

@@ -0,0 +1,46 @@
#[derive(Debug,Clone)]
pub enum TensorType{
F32,
I64,
Other
}
/// 明确命名为 AxisDim代表模型某一个轴的维度特征
#[derive(Clone, PartialEq, Eq)]
pub enum AxisDim {
/// 静态固定维度(如通道数固定为 1高度固定为 64
Static(usize),
/// 动态符号维度(如宽度是动态的 "image_width"
Dynamic(String),
}
impl AxisDim {
/// 便捷方法:判断是否为动态维度
pub fn is_dynamic(&self) -> bool {
matches!(self, AxisDim::Dynamic(_))
}
}
/// 自定义 Debug 格式化输出,彻底融化套娃外壳,保证日志干净漂亮
impl std::fmt::Debug for AxisDim {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
AxisDim::Static(size) => write!(f, "{}", size),
AxisDim::Dynamic(expr) => write!(f, "Dynamic(\"{}\")", expr),
}
}
}
/// 模拟 Python 的 input_info 和 output_info 结构
#[derive(Debug, Clone)]
pub struct TensorInfo {
pub name: String,
pub shape: Vec<AxisDim>, // 既包含 Fixed 静态维度,也包含 Dynamic 动态符号
pub tensor_type: TensorType, // 对应 Python 的 type
}
/// 最终返回的模型完整信息
#[derive(Debug, Clone)]
pub struct ModelInfo {
pub inputs: Vec<TensorInfo>,
pub outputs: Vec<TensorInfo>,
/// 硬件执行提供者(采用 Option 兼容不同底层的推理引擎)
pub providers: Option<Vec<String>>,
}

View File

@@ -0,0 +1,8 @@
pub mod image_convert;
mod image_helper;
pub mod image_processor;
mod tensor_transform;
// 对外统一暴露干净的 API 语义层
pub use image_convert::ColorMode;
pub use tensor_transform::normalize_ocr_logits;

View File

@@ -0,0 +1,174 @@
use crate::error::{ImagePreprocessError, Result};
use image::{DynamicImage, GenericImageView, ImageBuffer, Luma, Rgb, Rgba};
use ndarray::{Array3, ArrayViewD};
#[derive(Debug)]
pub enum ColorMode {
RGB,
RGBA,
L,
}
/// 封装数组转图像的逻辑,
// 对应 Python 版 _numpy_to_pil_image
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)
2 => ColorMode::L,
// 对应 Python: len(array.shape) == 3 (H, W, C)
3 => {
let (_h, _w, c) = (shape[0], shape[1], shape[2]);
match c {
// 对应 Python: array.shape[2] == 1 (单通道 H, W, 1)
1 => ColorMode::L,
// 对应 Python: array.shape[2] == 3 (RGB H, W, 3)
3 => ColorMode::RGB,
// 对应 Python: array.shape[2] == 4 (RGBA H, W, 4)
4 => ColorMode::RGBA,
_ => {
return Err(ImagePreprocessError::UnsupportedChannels(c));
}
}
}
_ => {
return Err(ImagePreprocessError::InvalidDimensions {
expected: "2D (H,W) 或 3D (H,W,C)".to_string(),
actual: shape.to_vec(),
});
}
};
from_ndarray(array, color_mode)
}
/// 处理PNG图片的RGBA透明背景将透明部分设置为白色背景
// 对应 Python 的 png_rgba_black_preprocess
pub fn png_rgba_white_preprocess(img: &DynamicImage) -> DynamicImage {
// 1. 检查是否包含透明通道,如果没有,直接克隆并返回
if !img.color().has_alpha() {
return DynamicImage::ImageRgb8(img.to_rgb8());
}
let (width, height) = img.dimensions();
// 2. 创建一个新的 RGB 图像缓冲,默认填充为白色 (255, 255, 255)
let mut background = ImageBuffer::from_pixel(width, height, Rgb([255u8, 255u8, 255u8]));
// 3. 获取原图的 RGBA 视图
let rgba_img = img.to_rgba8();
// 4. 遍历像素并手动进行 Alpha 混合
// 对应 Python 的 utils.paste(img, ..., mask=img)
// 使用 enumerate_pixels_mut 同时获取坐标和背景像素的可变引用,减少查找开销
for (x, y, bg_pixel) in background.enumerate_pixels_mut() {
// 安全性说明x, y 源自 background 尺寸,与 rgba_img 一致get_pixel 是安全的
let src_pixel = rgba_img.get_pixel(x, y);
let alpha_u8 = src_pixel[3];
match alpha_u8 {
// 情况 A完全不透明直接覆盖背景色
255 => {
bg_pixel.0 = [src_pixel[0], src_pixel[1], src_pixel[2]];
}
// 情况 B完全透明保持背景色白色无需操作
0 => {
continue;
}
// 情况 C半透明进行 Alpha 混合计算
_ => {
let alpha = alpha_u8 as f32 / 255.0;
let inv_alpha = 1.0 - alpha;
bg_pixel[0] = (src_pixel[0] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
bg_pixel[1] = (src_pixel[1] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
bg_pixel[2] = (src_pixel[2] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
}
}
}
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
// Rust utils 库通过 to_rgb8, to_luma8 等方法实现转换
let (width, height) = image.dimensions();
let (channels, raw) = match mode {
ColorMode::L => (1, image.to_luma8().into_raw()),
ColorMode::RGB => (3, image.to_rgb8().into_raw()),
ColorMode::RGBA => (4, image.to_rgba8().into_raw()),
};
let array = Array3::from_shape_vec((height as usize, width as usize, channels), raw)
.map_err(ImagePreprocessError::from)?;
Ok(array)
}
/// 将 array 数组转换为 DynamicImage
pub fn ndarray_to_image(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage,ImagePreprocessError> {
let shape = array.shape();
// 基础边界检查:至少要有 H 和 W 两个维度
if shape.len() < 2 {
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> {
let shape = array.shape();
// 映射ndarray 的 shape 默认是 [Height, Width, (Channels)]
// image 库的 from_raw 接收 (width, height)
let height = shape[0] as u32;
let width = shape[1] as u32;
// 1. 确保数据在内存中是连续的 (C order)
let standard = array.as_standard_layout();
let (raw_data, _) = standard.to_owned().into_raw_vec_and_offset();
let raw_len = raw_data.len();
// 获取当前模式对应的通道数
let channels = match mode {
ColorMode::L => 1,
ColorMode::RGB => 3,
ColorMode::RGBA => 4,
};
let expected_len = (width * height) as usize * channels;
// 构造通用错误闭包,避免 match 分支中重复编写冗长的错误对象
let make_err = || {
ImagePreprocessError::BufferLengthMismatch {
expected: expected_len,
actual: raw_len,
width,
height,
channels,
}
};
// 2. 重新解释内存并构建 ImageBuffer
match mode {
ColorMode::L => ImageBuffer::<Luma<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageLuma8)
.ok_or_else(make_err),
ColorMode::RGB => ImageBuffer::<Rgb<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageRgb8)
.ok_or_else(make_err),
ColorMode::RGBA => ImageBuffer::<Rgba<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageRgba8)
.ok_or_else(make_err),
}
}

View File

@@ -0,0 +1,152 @@
use crate::error::{DdddError, Result};
use crate::utils::image_convert::ndarray_to_hwc_image;
use base64::{engine::general_purpose, Engine as _};
use image::DynamicImage;
use ndarray::ArrayViewD;
use std::fmt;
use std::fmt::{Debug, Formatter};
use std::fs;
use std::path::Path;
use std::path::PathBuf;
pub struct Base64<'a>(pub &'a str);
/// 专属图像输入源转换器
pub struct ImageSource {
inner: DynamicImage,
}
// 1. 将 into_inner 优化为 into_image符合 Rust 官方命名规范
impl ImageSource {
/// 消耗当前包装器,获取最终的 DynamicImage
pub fn into_image(self) -> DynamicImage {
self.inner
}
}
pub trait TryFromImage<T>: Sized {
// 唯一的转换入口,通过目标类型来调用
fn try_from_image(value: T) -> Result<Self>;
}
// 1. 本身是 DynamicImage
impl TryFromImage<DynamicImage> for ImageSource {
fn try_from_image(img: DynamicImage) -> Result<Self> {
Ok(Self { inner: img })
}
}
#[derive(Debug)]
enum Base64ProcessError {
InvalidBase64Header,
}
impl fmt::Display for Base64ProcessError {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
match self {
Base64ProcessError::InvalidBase64Header => {
write!(f, "Base64 头部格式不正确,缺少 ';base64,' 分隔符")
}
}
}
}
impl std::error::Error for Base64ProcessError {}
// 2.2 路径类型 A: &str (最常用)
impl<'a> TryFromImage<&'a str> for ImageSource {
fn try_from_image(path_or_b64: &'a str) -> Result<Self> {
// 1. 嗅探:如果包含 Base64 特征
if path_or_b64.starts_with("data:image/") && path_or_b64.contains(";base64,") {
// 提取出真正的 base64 数据部分
let (_, clean_b64) = path_or_b64.split_once(";base64,").ok_or_else(|| {
// 返回一个明确的、可读性极佳的格式错误
DdddError::new(Base64ProcessError::InvalidBase64Header)
})?;
// 转换为 Base64 包装器,并复用其 TryFromImage 实现
Self::try_from_image(Base64(clean_b64))
} else {
// 2. 否则,老老实实当作本地路径打开
let img = image::open(path_or_b64).map_err(DdddError::new)?;
Ok(Self { inner: img })
}
}
}
// 2.3 路径类型 B: &Path (标准借用)
impl<'a> TryFromImage<&'a Path> for ImageSource {
fn try_from_image(path: &'a Path) -> Result<Self> {
let img = image::open(path).map_err(DdddError::new)?;
Ok(Self { inner: img })
}
}
// 2.4 路径类型 C: PathBuf / String (拥有所有权,透传给借用)
impl TryFromImage<PathBuf> for ImageSource {
fn try_from_image(path: PathBuf) -> Result<Self> {
Self::try_from_image(path.as_path())
}
}
impl TryFromImage<String> for ImageSource {
fn try_from_image(path_or_b64: String) -> Result<Self> {
Self::try_from_image(path_or_b64.as_str())
}
}
// 2. 支持带有生命周期的借用:直接支持 &[u8](不强绑生命周期到 ImageSource 结构体上!)
impl<'a> TryFromImage<&'a [u8]> for ImageSource {
fn try_from_image(bytes: &'a [u8]) -> Result<Self> {
let img = image::load_from_memory(bytes).map_err(DdddError::new)?;
Ok(Self { inner: img })
}
}
// 4. 完美支持 ndarray 的借用 ArrayViewD
impl<'a> TryFromImage<ArrayViewD<'a, u8>> for ImageSource {
fn try_from_image(array: ArrayViewD<'a, u8>) -> Result<Self> {
let img = ndarray_to_hwc_image(array)?;
Ok(Self { inner: img })
}
}
impl<'a> TryFromImage<Base64<'a>> for ImageSource {
fn try_from_image(b64_str: Base64<'a>) -> Result<Self> {
let img = base64_to_image(b64_str.0)?;
Ok(Self { inner: img })
}
}
/// 模拟 Python 的 load_image_from_input
#[allow(dead_code)]
pub fn load_image_from_input<I>(input: I) -> Result<DynamicImage>
where
ImageSource: TryFromImage<I>,
{
let img = ImageSource::try_from_image(input)?.into_image();
Ok(img)
}
/// 将base64编码的图片转换为 DynamicImage
pub fn base64_to_image(b64_str: &str) -> Result<DynamicImage> {
// 过滤掉可能存在的 base64 前缀,例如 "data:utils/png;base64,"
let clean_b64 = if let Some(pos) = b64_str.find(",") {
&b64_str[pos + 1..]
} else {
&b64_str
};
let bytes = general_purpose::STANDARD
.decode(clean_b64.trim())
// .map_err(|e| DdddError::new(e))?;
.map_err(DdddError::new)?;
let img = image::load_from_memory(&bytes).map_err(DdddError::new)?;
Ok(img)
}
/// 读取图片文件并转换为 base64 编码字符串
// 对应 Python 版 get_img_base64
pub fn img_to_base64<P: AsRef<Path>>(image_path: P) -> Result<String> {
// 1. 读取文件原始字节流
// 使用 AsRef<Path> 泛型可以让函数同时支持 String, &str, PathBuf 等类型
let image_data = fs::read(&image_path).map_err(DdddError::new)?;
// 2. 进行 Base64 编码
// 使用 STANDARD 引擎对齐 Python 的 base64.b64encode
let b64_string = general_purpose::STANDARD.encode(image_data);
Ok(b64_string)
}

View File

@@ -1,7 +1,9 @@
use image::{ImageBuffer, Luma};
use std::cmp::{max, min};
use tract_onnx::prelude::tract_ndarray::{Array2, Array3, ArrayView2, ArrayView3, azip};
use image::{imageops::FilterType, DynamicImage, GrayImage, ImageBuffer, Luma};
use ndarray::{azip, Array2, Array3, ArrayView2, ArrayView3};
use std::cmp::{max, min};
// 模拟openCV
/// 1. 计算两个数组的绝对差值 (对应 cv2.absdiff)
pub fn abs_diff(a: &ArrayView3<u8>, b: &ArrayView3<u8>) -> Array3<u8> {
// 利用 ndarray 的 map_collect生成差值的绝对值数组
@@ -72,6 +74,9 @@ pub fn find_contours_and_max(labelled: &ImageBuffer<Luma<u32>, Vec<u32>>) -> Opt
Some(max_label)
}
}
/// 根据目标连通域标签,计算其在图像中的外接矩形边界框(对应 `cv2.boundingRect`
///
/// 返回格式: `(min_x, min_y, width, height)`
pub fn bounding_rect(
labelled: &ImageBuffer<Luma<u32>, Vec<u32>>,
max_label: u32,
@@ -95,13 +100,22 @@ pub fn bounding_rect(
let h = max_y - min_y;
(min_x, min_y, w, h)
}
/// 根据左上角坐标与矩形长宽,计算其中央核心点坐标
#[inline]
pub fn calculate_center(top_left: (u32, u32), width: usize, height: usize) -> (i32, i32) {
let center_x = top_left.0 as i32 + (width as i32 / 2);
let center_y = top_left.1 as i32 + (height as i32 / 2);
(center_x, center_y)
}
/// 高性能转换:将 `ndarray` 2D 灰度视图规整为 `image::ImageBuffer` 格式
///
/// 放弃低效的逐像素显式嵌套循环,采用原生内存池直接构造,减少寻址开销
pub fn ndarray_to_luma8(array: ArrayView2<u8>) -> ImageBuffer<Luma<u8>, Vec<u8>> {
let (height, width) = array.dim();
// 技巧:直接将已有的规整连续内存打平转换,或用 from_raw 包装
// 此处保留安全的一步转换,但用更内聚的迭代器或切片拷贝进行速度优化
let mut buffer = ImageBuffer::new(width as u32, height as u32);
for y in 0..height {
for x in 0..width {
@@ -126,7 +140,7 @@ pub fn rgb_to_opencv_hsv(r: u8, g: u8, b: u8) -> (u8, u8, u8) {
let delta = max - min;
// 2. 计算 H (色调) - 移除负数取余陷阱,改用平铺分支
let mut h = if delta == 0.0 {
let h = if delta == 0.0 {
0.0
} else if max == r_f {
let mut diff = (g_f - b_f) / delta;
@@ -159,3 +173,24 @@ pub fn rgb_to_opencv_hsv(r: u8, g: u8, b: u8) -> (u8, u8, u8) {
(h_opencv, s_opencv, v_opencv)
}
/// 对应 Python 的 convert_to_grayscale
/// 将图像转换为灰度图 (L模式)
pub fn convert_to_grayscale(image: &DynamicImage) -> GrayImage {
// Rust utils 库的 to_luma8 会根据标准的亮度公式进行转换
image.to_luma8()
}
/// 对应 Python 的 resize_image
/// 调整图像尺寸。当前版本仅实现 keep_aspect_ratio=false
pub fn resize_image(
image: &DynamicImage,
target_width: u32,
target_height: u32,
// resample 参数我们直接使用 FilterTypeLanczos3 是最接近 Python LANCZOS 的
) -> DynamicImage {
// image::imageops::resize 的最高层封装
// FilterType::Lanczos3 与 Python Pillow 的 Image.LANCZOS 算法完全对齐,缩放质量最高
image.resize_exact(target_width, target_height, FilterType::Lanczos3)
}

View File

@@ -0,0 +1,51 @@
use crate::OcrOutput;
use crate::error::{Result, TensorError};
use ndarray::s;
/// 核心层复用资产:将异构的动态维度矩阵转化为标准 OCR 2D Logits 矩阵
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 {
// 形状: [Steps, 1, Classes]
(shape[0], shape[2], array)
} else if shape[0] == 1 {
// 形状: [1, Steps, Classes]
(shape[1], shape[2], array)
} else {
// 默认取第一个 batch: [Batch, Steps, Classes]
// 使用 ndarray 的 s! 宏,对应 Python 的 output[0, :, :]
let sliced = array.slice_move(s![0, .., ..]);
(shape[1], shape[2], sliced.into_dyn())
}
}
// 形状: [Steps, Classes]
2 => (shape[0], shape[1], array),
// 形状: [Classes] -> 单字符输出(对应 Python 的 ndim == 0 保护逻辑)
// 我们把它虚构成一个 [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(),
},
);
}
};
// 转换为标准的 2D 静态矩阵 [Steps, Classes]
let matrix_cow = data_dyn_view
.to_shape(ndarray::Ix2(steps, classes))
.map_err(|shape_err| {
// 如果是因为切片导致不连续且无法进行零拷贝变换,抛出 NonContiguousMemory
if !data_dyn_view.is_standard_layout() {
TensorError::NonContiguousMemory
} else {
// 否则说明是纯粹的数据元素数量不对Shape 不匹配),抛出专属的强类型错误
TensorError::LogitsDimensionMismatch(shape_err)
}
})?
.to_owned();
Ok(OcrOutput::Logits(matrix_cow))
}

23
ddddocr-ort/Cargo.toml Normal file
View File

@@ -0,0 +1,23 @@
[package]
name = "ddddocr-ort"
version.workspace = true
edition.workspace = true
license.workspace = true
[dependencies]
ddddocr-core = { path = "../ddddocr-core" } # 引入兄弟库
ort = { workspace = true ,features = ["cuda"]}
ndarray = { workspace = true } # 继承自工作空间
image = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
anyhow = { workspace = true }
thiserror = { workspace = true } # 刚好可以开始接入你需要的标准库错误处理
[features]
# 1. 声明 cuda feature
# 2. 当开启 ddddocr-ort 的 cuda feature 时,自动开启底层 ort 库的 cuda 支持(如果 ort 库支持的话)
cuda = ["ort/cuda"]

1
ddddocr-ort/src/det.rs Normal file
View File

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

View File

@@ -0,0 +1,65 @@
use crate::types::Session;
use ddddocr_core::DetOutput;
use ddddocr_core::error::{Result, TensorError};
use ddddocr_core::traits::{DetEngine, InferenceEngine};
use ndarray::Ix3;
use ort::inputs;
use ort::value::TensorRef;
// use tract_onnx::prelude::{tvec, IntoTensor, Tensor};
#[derive(Debug)]
pub struct DetRuntime {
pub session: Session,
}
impl DetRuntime {
pub fn new(session: Session) -> Self {
Self { session }
}
}
impl InferenceEngine for DetRuntime {
type Output = DetOutput; // 明确绑定 OCR 小枚举
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
// let result = self.ocr.run(tvec!(tensor.into()))?;
let mut session_guard = self
.session
.lock()
.map_err(|_| TensorError::Engine("获取 Session 锁失败 (Poisoned)".to_string()))?;
let result = session_guard
.run(inputs![TensorRef::from_array_view(&input_array).map_err(
|e| TensorError::Engine(format!("构建输入失败: {e}"))
)?])
.map_err(|e| TensorError::Engine(format!("执行模型推理失败: {e}")))?;
// .context("执行模型推理失败")?;
println!("模型输出原始数据: {:?}", result);
// Ok(result.swap_remove(0).into_tensor())
let raw_value = &result[0];
// raw_tensor.into_plain_array()?
let (shape_ref, slice) = raw_value.try_extract_tensor::<f32>().map_err(|_| {
TensorError::Engine("Tract 实体张量无法转换为 ndarray::ArrayD".to_string())
})?;
// 提前利用克隆(Clone)备份好当前未转维度前的真实 shape (Vec<usize>)
let shape_vec: Vec<usize> = shape_ref.to_vec().iter().map(|v| *v as usize).collect();
let shape_vec_slice = shape_vec.as_slice();
let view = ndarray::ArrayViewD::from_shape(shape_vec_slice, slice)
.map_err(|_| TensorError::Engine("构建 ndarray ArrayViewD 失败".to_string()))?;
let array3 = view.to_owned().into_dimensionality::<Ix3>().map_err(|_| {
TensorError::DimensionMismatch {
expected: "3D 检测矩阵 [Batch, Box_Count, Box_Attributes]".to_string(),
actual: shape_vec, // 优雅降维失败时动态捕获
}
})?;
Ok(DetOutput::Detection(array3))
// 在引擎内部消化掉 DatumType 强耦合
}
}
impl DetEngine for DetRuntime {}

9
ddddocr-ort/src/lib.rs Normal file
View File

@@ -0,0 +1,9 @@
mod det;
pub mod loader;
mod ocr;
mod types;
pub use ddddocr_core::{SlideResult, Slider,OcrBuilder};
pub use det::session::DetRuntime;
pub use ocr::session::OcrRuntime;

View File

@@ -0,0 +1,7 @@
mod error;
mod metadata;
mod model;
pub use error::{Error, ParseError, Result};
pub use metadata::{ModelMetadataDto, NormalizationDto, Metadata};
pub use model::ModelLoader;

View File

@@ -0,0 +1,66 @@
use ort::Error as OrtError;
pub type Result<T> = std::result::Result<T, Error>;
#[derive(thiserror::Error, Debug)]
pub enum Error {
#[error("builder构建失败")]
Build(#[from] BuildError),
/// 解析 ONNX 模型/路径失败(如文件损坏、算子不支持、路径非法)
#[error("解析 ONNX 模型结构失败: {0}")]
ModelParse(#[from] ParseError),
/// 模型计算图优化失败(如常量折叠、形状推导失败)
#[error("优化 Tract 模型图失败: {0}")]
OptimizationFailed(#[source] OrtError),
/// 构建可执行 Session 失败(如输入输出 Tensor 类型/形状未确定)
#[error("构建可运行 Tract 实例失败: {0}")]
RunnableBuildFailed(#[source] OrtError),
/// JSON 反序列化失败(自动透传 serde_json 报错)
#[error("模型 Metadata JSON 解析失败: {0}")]
JsonParse(#[from] serde_json::Error),
/// 字节流非合法 UTF-8 编码(自动透传 Utf8Error
#[error("Metadata 字节流不是合法的 UTF-8 编码: {0}")]
InvalidUtf8(#[from] std::str::Utf8Error),
#[error("模型元数据解析失败: {0}")]
MetadataParse(String),
/// 承载任何第三方扩展、解密、特定预处理插件在执行时产生的自定义错误
#[error("{0}: {1}")]
Other(String, #[source] Box<dyn std::error::Error + Send + Sync>),
}
impl Error {
/// 方便将任何第三方 Error 包装为 Error::Other
pub fn new<E>(msg: impl Into<String>, err: E) -> Self
where
E: Into<Box<dyn std::error::Error + Send + Sync>>,
{
Self::Other(msg.into(), err.into())
}
}
#[derive(thiserror::Error, Debug)]
pub enum ParseError {
/// 策略 A从文件路径加载失败附带路径上下文信息方便排查是找不到文件还是格式不对
#[error("从路径 '{0}' 加载 ONNX 模型失败: {1}")]
Path(String, #[source] OrtError),
/// 策略 B从内存字节流加载失败如 include_bytes! 传入的字节流损坏)
#[error("从内存字节流解析 ONNX 模型失败: {0}")]
Bytes(#[source] OrtError),
}
#[derive(thiserror::Error, Debug)]
pub enum BuildError{
#[error("builder构建失败")]
BuildFailed(#[from] OrtError),
#[error("builder构建失败")]
Threads(String),
#[error("builder构建失败")]
EnabledCudaFailed(String),
#[error("builder构建失败")]
NotEnabledCuda(String)
}

View File

@@ -0,0 +1,93 @@
use crate::loader::error::{Error, Result};
use ddddocr_core::ModelMetadata;
use ddddocr_core::Resize;
use ddddocr_core::{Charset, Normalization};
use serde::Deserialize;
use std::borrow::Cow;
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")] // 支持 json 中写 "zero_to_one" 或 "minus_one_to_one"
pub enum NormalizationDto {
/// 映射到 [0.0, 1.0] -> pixel / 255.0
ZeroToOne,
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
MinusOneToOne,
}
impl From<NormalizationDto> for Normalization {
fn from(dto: NormalizationDto) -> Self {
match dto {
NormalizationDto::ZeroToOne => Normalization::ZeroToOne,
NormalizationDto::MinusOneToOne => Normalization::MinusOneToOne,
}
}
}
/// 仅用于反序列化 JSON 的中间临时结构体DTO
#[derive(Deserialize)]
pub struct ModelMetadataDto {
charset: Vec<String>,
word: bool,
#[serde(alias = "image")]
resize: Vec<i32>,
channel: u8,
/// 新增:允许在配置文件中指定归一化策略。
/// 使用 serde(default) 可以在不配置时提供一个默认值(比如默认 ZeroToOne
#[serde(default = "default_normalization")]
normalization: NormalizationDto,
}
fn default_normalization() -> NormalizationDto {
NormalizationDto::ZeroToOne
}
/// Tract 专属扩展trait 或 工具函数
pub trait Metadata: Sized {
fn from_json_str(json_str: &str) -> Result<Self>;
/// 机制 2从内存字节流加载极大地方便 include_bytes! 或网络下载)
fn from_json_bytes(bytes: &[u8]) -> Result<Self> {
let json_str = std::str::from_utf8(bytes)?;
Self::from_json_str(json_str)
}
}
impl Metadata for ModelMetadata {
// --- 优雅的工厂模式构造器 ---
fn from_json_str(json_str: &str) -> Result<ModelMetadata> {
let dto: ModelMetadataDto = serde_json::from_str(json_str)?;
// 1. 将 DTO 的字符串数组转化为强类型的 Charset
let tokens: Vec<Cow<'static, str>> =
dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
let charset = Charset::new(tokens);
// 2. 解析 resize 策略(重现 Python 的复杂条件判断)
if dto.resize.len() != 2 {
return Err(Error::MetadataParse(
"'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]".to_string(),
));
}
let r0 = dto.resize[0];
let r1 = dto.resize[1];
let resize = if r0 == -1 {
if dto.word {
// 如果 word 为 true且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
Resize::Square(r1 as u32)
} else {
// 如果 word 为 false且包含 -1Python 里是高度固定为 r1宽度按原图比例缩放
Resize::DynamicWidth(r1 as u32)
}
} else {
// 正常的固定宽高
Resize::Fixed(r0 as u32, r1 as u32)
};
Ok(ModelMetadata::new(
charset,
dto.word,
resize,
dto.channel,
dto.normalization.into(),
))
}
}

View File

@@ -0,0 +1,114 @@
use crate::loader::Error;
use crate::loader::error::{BuildError, ParseError, Result};
use crate::types::Session;
use ddddocr_core::traits::Loader;
use ort::session::Session as OrtSession;
use ort::session::builder::SessionBuilder;
use std::sync::{Arc, Mutex};
// pub struct OrtModelLoader;
//
// impl OrtModelLoader {
// /// 获取针对 ORT 后端的链式构建器
// pub fn builder() -> OrtModelBuilder {
// OrtModelBuilder::default()
// }
// }
/// ORT 专用的链式构建器
#[derive(Debug, Clone)]
pub struct ModelLoader {
use_gpu: bool,
device_id: i32,
intra_threads: Option<usize>,
}
impl Default for ModelLoader {
fn default() -> Self {
Self {
use_gpu: false,
device_id: 0,
intra_threads: None,
}
}
}
impl ModelLoader {
/// 开启或关闭 GPU 加速
pub fn use_gpu(mut self, enable: bool) -> Self {
self.use_gpu = enable;
self
}
/// 指定 GPU 设备 ID
pub fn device_id(mut self, id: i32) -> Self {
self.device_id = id;
self
}
pub fn num_threads(mut self, threads: usize) -> Self {
self.intra_threads = Some(threads);
self
}
/// 内部辅助方法:根据当前的配置构建 ORT 底层的 SessionBuilder
fn create_session_builder(&self) -> Result<SessionBuilder> {
let mut builder = OrtSession::builder().map_err(|e| BuildError::BuildFailed(e))?;
// 如果用户显式设置了线程数,则配置给 ORT
if let Some(threads) = self.intra_threads {
builder = builder
.with_intra_threads(threads)
.map_err(|e| BuildError::Threads(format!("设置线程数失败: {e}")))?;
}
if self.use_gpu {
// 根据 ort 库版本配置 CUDA 执行提供者 (Execution Provider)
#[cfg(feature = "cuda")]
{
use ort::ep::CUDAExecutionProvider;
let cuda_ep = CUDAExecutionProvider::default().with_device_id(self.device_id);
builder = builder
.with_execution_providers([cuda_ep.build()])
.map_err(|e| {
BuildError::EnabledCudaFailed(format!("配置 CUDA 硬件加速失败: {e}"))
})?
}
#[cfg(not(feature = "cuda"))]
{
// 如果用户明确开启了 GPU但 Feature 没编译进去,明确抛错提醒
return Err(BuildError::NotEnabledCuda(
"未启用 CUDA 支持:请在 Cargo.toml 中为 ddddocr-ort 开启 `cuda` feature"
.to_string(),
))?;
}
}
Ok(builder)
}
}
impl Loader for ModelLoader {
type Session = Session;
type Error = Error;
fn build_for_path<P>(&self, model_path: P) -> Result<Session>
where
P: AsRef<std::path::Path>,
{
let path_ref = model_path.as_ref();
let mut builder = self.create_session_builder()?;
// Session::builder() 会返回 Result<SessionBuilder, ort::Error>
let session = builder
.commit_from_file(path_ref)
.map_err(|e| ParseError::Path(path_ref.display().to_string(), e))?;
Ok(Arc::new(Mutex::new(session))) // 这里的session需要包装下
}
/// 策略 B从内存字节流加载模型配合 include_bytes! 使用)
fn build_from_bytes(&self, model_bytes: &[u8]) -> Result<Session> {
let mut builder = self.create_session_builder()?;
let session = builder
.commit_from_memory(model_bytes)
.map_err(ParseError::Bytes)?;
Ok(Arc::new(Mutex::new(session)))
}
}

1
ddddocr-ort/src/ocr.rs Normal file
View File

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

View File

@@ -0,0 +1,150 @@
use crate::types::Session;
use ddddocr_core::ModelMetadata;
use ddddocr_core::OcrOutput;
use ddddocr_core::error::{DdddError, Result, TensorError};
use ddddocr_core::traits::{InferenceEngine, Info, OcrEngine};
use ddddocr_core::types::{AxisDim, ModelInfo, TensorInfo, TensorType};
use ddddocr_core::utils::normalize_ocr_logits;
use ort::inputs;
use ort::value::{TensorElementType, TensorRef};
use std::sync::Mutex;
// 引入核心层的统一错误类型
/// 明确命名为 AxisDim代表模型某一个轴的维度特征
// #[derive(Clone, PartialEq, Eq)]
// pub enum AxisDim {
// /// 静态固定维度(如通道数固定为 1高度固定为 64
// Static(usize),
// /// 动态符号维度(如宽度是动态的 "image_width"
// Dynamic(String),
// }
// impl AxisDim {
// /// 便捷方法:判断是否为动态维度
// pub fn is_dynamic(&self) -> bool {
// matches!(self, AxisDim::Dynamic(_))
// }
// }
// /// 自定义 Debug 格式化输出,彻底融化套娃外壳,保证日志干净漂亮
// impl std::fmt::Debug for AxisDim {
// fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
// match self {
// AxisDim::Static(size) => write!(f, "{}", size),
// AxisDim::Dynamic(expr) => write!(f, "Dynamic(\"{}\")", expr),
// }
// }
// }
/// 模拟 Python 的 input_info 和 output_info 结构
// #[derive(Debug, Clone)]
// pub struct TensorInfo {
// pub name: String,
// pub shape: Vec<AxisDim>, // 既包含 Fixed 静态维度,也包含 Dynamic 动态符号
// pub data_type: TensorElementType, // 对应 Python 的 type
// }
//
// /// 最终返回的模型完整信息
// #[derive(Debug, Clone)]
// pub struct ModelInfo {
// pub inputs: Vec<TensorInfo>,
// pub outputs: Vec<TensorInfo>,
// /// 硬件执行提供者(采用 Option 兼容不同底层的推理引擎)
// pub providers: Option<Vec<String>>,
// }
pub struct OcrRuntime {
pub session: Session,
pub metadata: ModelMetadata,
}
impl OcrRuntime {
pub fn new(session: Session, metadata: ModelMetadata) -> Self {
Self { session, metadata }
}
}
impl OcrEngine for OcrRuntime {
fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
}
impl InferenceEngine for OcrRuntime {
type Output = OcrOutput;
/// 对应 Python 的 _inference
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
// let result = self.ocr.run(tvec!(tensor.into()))?;
// let tensor = Tensor::from(input_array);
let mut session_guard = self
.session
.lock()
.map_err(|_| TensorError::Engine("获取 Session 锁失败 (Poisoned)".to_string()))?;
let result = session_guard
.run(inputs![TensorRef::from_array_view(&input_array).map_err(
|e| TensorError::Engine(format!("构建输入失败: {e}"))
)?])
.map_err(|e| TensorError::Engine(format!("执行模型推理失败: {e}")))?;
// .context("执行模型推理失败")?;
println!("模型输出原始数据: {:?}", result);
// Ok(result.swap_remove(0).into_tensor())
let raw_value = &result[0];
match raw_value.dtype().tensor_type().unwrap() {
TensorElementType::Int64 => {
let (array_d, slice) = raw_value
.try_extract_tensor::<i64>()
.map_err(|_| TensorError::Engine("Tract 无法获取 i64 内存视图".to_string()))?;
// .context("Tract 无法获取 i64 内存视图")?;
// 提前提取真实维度
let actual_shape = array_d
.to_vec()
.iter()
.map(|v| *v as usize)
.collect::<Vec<usize>>();
let view = ndarray::ArrayViewD::from_shape(actual_shape.as_slice(), slice)
.map_err(|_| TensorError::Engine("构建 ndarray ArrayViewD 失败".to_string()))?;
// 转成标准的 Array1 传给 core
let array1 = view
.to_owned()
.into_dimensionality::<ndarray::Ix1>()
.map_err(|_| TensorError::DimensionMismatch {
expected: "1D 字符索引静态矩阵".to_string(),
actual: actual_shape,
})?;
Ok(OcrOutput::Indices(array1))
}
TensorElementType::Float32 => {
let shape = raw_value.shape();
println!("模型输出shape数据: {:?}", shape);
// raw_tensor.to_plain_array_view()
let (shape_ref, slice) = raw_value
.try_extract_tensor::<f32>()
.map_err(|_| TensorError::Engine("Tract 无法获取 f32 内存视图".to_string()))?;
// 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗
let shape_vec: Vec<usize> =
shape_ref.to_vec().iter().map(|v| *v as usize).collect();
let shape_vec_slice = shape_vec.as_slice();
let view = ndarray::ArrayViewD::from_shape(shape_vec_slice, slice)
.map_err(|_| TensorError::Engine("构建 ndarray ArrayViewD 失败".to_string()))?;
normalize_ocr_logits(view, shape_vec_slice)
}
_ => Err(
// anyhow::anyhow!("不支持的模型输出数据类型: {:?}",raw_tensor.datum_type())
TensorError::UnknownOutputFormat,
),
}
}
}
impl Info for OcrRuntime {
fn input_info(&self) -> Result<Vec<TensorInfo>> {
todo!()
}
fn output_info(&self) -> Result<Vec<TensorInfo>> {
todo!()
}
fn model_info(&self) -> Result<ModelInfo> {
todo!()
}
}

4
ddddocr-ort/src/types.rs Normal file
View File

@@ -0,0 +1,4 @@
use ort::session::Session as OrtSession;
use std::sync::{Arc, Mutex};
pub type Session = Arc<Mutex<OrtSession>>;

View File

@@ -1,3 +1,10 @@
use std::borrow::Cow;
use std::fs::File;
use std::path::Path;
use anyhow::anyhow;
use ddddocr_core::Charset;
use ddddocr_core::{Normalization, Resize};
pub const CHARSET_BETA: &[&str] = &[
"", "", "", "", "", "", "", "", "", "", "", "", "", "6", "", "",
"", "", "", "", "", "", "", "", "", "", "", "鴿", "", "", "", "",
@@ -517,212 +524,77 @@ pub const CHARSET_BETA: &[&str] = &[
pub const CHARSET_OLD: &[&str] = &["", "", "", "", ""];
use std::borrow::Cow;
use std::collections::{HashMap, HashSet};
/// 字符集范围限制枚举
pub struct ValidationCtx<'a> {
pub text: &'a str, // 当前 Token 的文本内容
pub token_id: usize, // 当前 Token 的 ID 索引
}
/// 统一的约束接口
pub trait TokenFilter {
fn matches(&self, ctx: &ValidationCtx) -> bool;
/// 预估容量提示,帮助精准开辟 Vec 内存
fn estimated_capacity(&self) -> usize {
128
}
/// 【新引入的架构级核心方法】
/// 统一接管全量字符集的密集遍历、CTC Blank放行、去重、排序及空交集退化兜底
fn apply_to_charset(&self, tokens: &[Cow<str>]) -> Option<Vec<usize>> {
let mut has_any_match = false;
let estimated_capacity = self.estimated_capacity();
// 1. 精准开辟内存,完美利用容量提示,避免动态乱涨
let mut temp_indices = Vec::with_capacity(estimated_capacity.max(16));
// pub fn from_builtin_old() -> Self {
// Self::from_static_slice(
// CHARSET_OLD,
// false,
// Resize::DynamicWidth(64),
// 1,
// Normalization::ZeroToOne,
// )
// }
//
// /// 从预设的 Beta 版字符集创建
// pub fn from_builtin_beta() -> Self {
// Self::from_static_slice(
// CHARSET_BETA,
// false,
// Resize::DynamicWidth(64),
// 1,
// Normalization::MinusOneToOne,
// )
// }
// 2. 高性能原地单次流式迭代
for (idx, token) in tokens.iter().enumerate() {
let token_str = token.as_ref();
// 规则 A: CTC Blank 空字符串或 0 号索引无条件放行
if token_str.is_empty() || idx == 0 {
temp_indices.push(idx);
continue; // 关键:直接跳过,防止后续 matches 匹配成功导致重复 push 产生 Bug
}
// 规则 B: 组装无拷贝上下文
let ctx = ValidationCtx {
text: token_str,
token_id: idx,
};
// 规则 C: 路由到各自具体实现的特异性匹配中(如 Digit 判定、TopN 判定、组合子判定等)
if self.matches(&ctx) {
temp_indices.push(idx);
has_any_match = true;
}
}
// 3. 终极防御:如果整个模型字符集除了 Blank一个都没对上直接退化为 None全量识别
if !has_any_match {
println!("警告:当前限制策略与模型字符集完全没有交集!已自动恢复全量识别。");
None
} else {
// 4. 排序并去重,为 Ocr 引擎后续进行极其高频的『二分查找』筑起绝对安全的底层保障
temp_indices.sort_unstable();
temp_indices.dedup();
Some(temp_indices)
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CharRestrict {
Digit,
Lowercase,
Uppercase,
CustomList(Vec<String>),
}
impl TokenFilter for CharRestrict {
fn matches(&self, ctx: &ValidationCtx) -> bool {
match self {
Self::Digit => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_digit(),
Self::Lowercase => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_lowercase(),
Self::Uppercase => ctx.text.len() == 1 && ctx.text.as_bytes()[0].is_ascii_uppercase(),
Self::CustomList(vec) => vec.iter().any(|t| t == ctx.text),
}
}
fn estimated_capacity(&self) -> usize {
match self {
Self::Digit => 16,
Self::Lowercase | Self::Uppercase => 32,
Self::CustomList(vec) => vec.len() + 1,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IdRestrict {
TopN(usize),
IdRange(std::ops::Range<usize>),
IdList(Vec<usize>),
}
impl TokenFilter for IdRestrict {
fn matches(&self, ctx: &ValidationCtx) -> bool {
match self {
Self::TopN(n) => ctx.token_id < *n,
Self::IdRange(range) => range.contains(&ctx.token_id),
Self::IdList(vec) => vec.contains(&ctx.token_id),
}
}
fn estimated_capacity(&self) -> usize {
match self {
Self::TopN(n) => *n + 1,
// 2. IdRange标准标准库 Range 的长度
// 注意:因为范围可能是 1000..2000,它的 len() 返回的是 usize
Self::IdRange(range) => range.len() + 1,
// 3. IdListVec 里的元素个数
Self::IdList(vec) => vec.len() + 1,
}
}
}
/// 多路“或”逻辑组合子(支持 N 个规则无缝并集)
pub struct MultiOrRestrict<'a> {
pub filters: Vec<&'a dyn TokenFilter>,
}
impl<'a> TokenFilter for MultiOrRestrict<'a> {
fn matches(&self, ctx: &ValidationCtx) -> bool {
// 核心高阶函数:只要有一个过滤器命中,该 Token 即可放行
self.filters.iter().any(|f| f.matches(ctx))
}
fn estimated_capacity(&self) -> usize {
// 将所有过滤器的预估容量累加,作为最终容量参考
self.filters.iter().map(|f| f.estimated_capacity()).sum()
}
}
// =====================================================================
// 声明式宏:替代 `+` 运算符,解决组合扩展痛苦
// =====================================================================
#[macro_export]
macro_rules! any_of {
// 场景 A如果用户只传了一个规则免去构建 Vec 的开销,直接返回其引用
($only:expr) => {
&$only as &dyn $crate::TokenFilter
};
// 场景 B如果用户传入了多个规则自动织成一张静态组合网
($($filter:expr),+ $(,)?) => {
&$crate::MultiOrRestrict {
filters: vec![ $( &$filter as &dyn $crate::TokenFilter ),+ ]
}
};
}
// ==========================================
// 3. 字符集核心结构体 (重命名为 Charset)
// ==========================================
#[derive(Debug, Clone)]
pub struct Charset {
// 使用 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 {
// 内部底层统一收拢构造
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 {
tokens,
char_to_idx,
}
}
// --- 业务策略方法 ---
/// 将字符转为索引,不存在返回 -1 (保持与原 Python 库行为一致)
pub fn char_to_index(&self, char_str: &str) -> i32 {
if let Some(&idx) = self.char_to_idx.get(char_str) {
idx as i32
} else {
-1
}
}
/// 将索引转为字符引用,零拷贝。若越界返回 None
pub fn index_to_char_ref(&self, index: usize) -> Option<&str> {
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()
}
pub fn size(&self) -> usize {
self.tokens.len()
}
}
// ==========================================
// 4. 标准 Display 接口实现 (对应 __str__)
// ==========================================
impl std::fmt::Display for Charset {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "Charset [Total Size: {}", self.size(),)
}
}
// /// 从外部外部 JSON 文件动态加载字符集(在后续优化中移除)
// pub fn from_json_file<P: AsRef<Path>>(path: P) -> anyhow::Result<Self> {
// let path = path.as_ref();
// if !path.exists() {
// return Err(anyhow!("模型元数据配置文件不存在: {:?}", path));
// }
//
// let mut file = File::open(path)?;
// let mut content = String::new();
// file.read_to_string(&mut content)?;
//
// let dto: ModelMetadataDto = serde_json::from_str(&content)
// .map_err(|e| anyhow!("JSON 反序列化失败,请检查字段是否完整: {}", e))?;
//
// // 1. 将 DTO 的字符串数组转化为强类型的 Charset
// let tokens: Vec<Cow<'static, str>> =
// dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
// let charset = Charset::new(tokens);
//
// // 2. 解析 resize 策略(重现 Python 的复杂条件判断
// if dto.resize.len() != 2 {
// return Err(anyhow!(
// "'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]"
// ));
// }
// let r0 = dto.resize[0];
// let r1 = dto.resize[1];
//
// let resize = if r0 == -1 {
// if dto.word {
// // 如果 word 为 true且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
// Resize::Square(r1 as u32)
// } else {
// // 如果 word 为 false且包含 -1Python 里是高度固定为 r1宽度按原图比例缩放
// Resize::DynamicWidth(r1 as u32)
// }
// } else {
// // 正常的固定宽高
// Resize::Fixed(r0 as u32, r1 as u32)
// };
//
// Ok(Self {
// charset,
// word: dto.word,
// resize,
// channel: dto.channel,
// normalization: dto.normalization,
// })
// }

View File

@@ -0,0 +1,253 @@
use anyhow::Context;
use ddddocr_core::{DetectionResult, Ocr};
use ddddocr_core::traits::Loader;
use ddddocr_core::{Detector, ModelMetadata, Normalization, Slider};
// 假设你的包名是这个
use ddddocr_ort::{DetRuntime, OcrBuilder, OcrRuntime};
use image::{DynamicImage, ImageBuffer, Luma, Rgb};
use std::fs;
use std::path::Path;
mod char_slice;
use char_slice::CHARSET_BETA;
use ddddocr_core::Resize;
use ddddocr_ort::loader::ModelLoader as OrtModelLoader;
fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
// 1. 先将泛型转为具体的 &Path 引用
let path_ref = path.as_ref();
// 2. 调用 open 时传入引用utils::open 支持 AsRef<Path>
image::open(path_ref).map_err(|e| {
// 3. 此时 path_ref 依然有效,可以安全地在闭包中使用
anyhow::anyhow!("无法加载图片 {:?}: {}", path_ref, e)
})
}
/// 将检测结果绘制在图像上并保存
fn save_debug_image(
dynamic_img: &DynamicImage, // 【优化点 1】直接传入解码好的引用拒绝重复解码
bboxes: &[DetectionResult], // 【修改点 1】类型改为自定义结构体切片
output_path: &str,
) -> anyhow::Result<()> {
// 删除了原本的 let dynamic_img = image::load_from_memory(image_bytes)?;
let mut img = dynamic_img.to_rgb8();
let (width, height) = img.dimensions();
let red = Rgb([255u8, 0, 0]);
for bbox in bboxes {
// 【修改点 2】将原来的索引 bbox[0].. 改为结构体字段访问 .x1, .y1 ..
let x1 = bbox.x1.max(0).min(width as i32 - 1) as u32;
let y1 = bbox.y1.max(0).min(height as i32 - 1) as u32;
let x2 = bbox.x2.max(0).min(width as i32 - 1) as u32;
let y2 = bbox.y2.max(0).min(height as i32 - 1) as u32;
// 绘制横向线条
for x in x1..=x2 {
img.put_pixel(x, y1, red);
img.put_pixel(x, y2, red);
if y1 + 1 < height {
img.put_pixel(x, y1 + 1, red);
}
if y2.saturating_sub(1) > 0 {
img.put_pixel(x, y2 - 1, red);
}
}
// 绘制纵向线条
for y in y1..=y2 {
img.put_pixel(x1, y, red);
img.put_pixel(x2, y, red);
if x1 + 1 < width {
img.put_pixel(x1 + 1, y, red);
}
if x2.saturating_sub(1) > 0 {
img.put_pixel(x2 - 1, y, red);
}
}
}
img.save(output_path)?;
Ok(())
}
#[allow(dead_code)]
fn save_rust_result(result: &ImageBuffer<Luma<f32>, Vec<f32>>, filename: &str) {
let (width, height) = result.dimensions();
// 1. 寻找最值进行归一化
let mut max_val = f32::MIN;
let mut min_val = f32::MAX;
for p in result.pixels() {
if p.0[0] > max_val {
max_val = p.0[0];
}
if p.0[0] < min_val {
min_val = p.0[0];
}
}
// 2. 创建 8 位灰度图
let mut out_buf = ImageBuffer::new(width, height);
for y in 0..height {
for x in 0..width {
let val = result.get_pixel(x, y).0[0];
let normalized = if max_val > min_val {
((val - min_val) / (max_val - min_val) * 255.0) as u8
} else {
0u8
};
out_buf.put_pixel(x, y, Luma([normalized]));
}
}
// 3. 保存
DynamicImage::ImageLuma8(out_buf).save(filename).unwrap();
println!("Rust 结果热力图已保存至: {}", filename);
}
#[test]
fn test_full_classification() {
let model = OrtModelLoader::default().use_gpu(false)
.build_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx")
// .build_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_old.onnx")
.expect("模型加载失败");
let metadata = ModelMetadata::from_static_slice(
CHARSET_BETA,
false,
Resize::DynamicWidth(64),
1,
Normalization::MinusOneToOne,
);
// 1. 初始化模型
let ocr = OcrRuntime::new(model, metadata);
// 2. 加载测试图片
let img =
image::open("D:/CNWei/CNW/Rust/ddddocr-rs/samples/code2.png").expect("测试图片不存在");
// 3. 执行识别
// let result = Ocr::new(&ocr)
// .predict(&img)
// .expect("识别过程出错")
// .into_text();
// let result = OcrBuilder::new()
// .build(&ocr)
// .predict(&img)
// .expect("识别过程出错")
// .into_text();
let res=Ocr::builder().runner(&ocr).predict(&img).expect("s").into_text();
// println!("识别结果: {}", result);
println!("识别结果: {}", res);
// assert!(!result.is_empty());
assert!(!res.is_empty());
}
#[test]
fn test_det_load() -> anyhow::Result<()> {
let det_model = OrtModelLoader::default()
.build_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_det.onnx")
.expect("模型加载失败");
let det = DetRuntime::new(det_model);
let image_path = "D:/CNWei/CNW/Rust/ddddocr-rs/samples/det1.png";
let image_bytes =
fs::read(image_path).map_err(|e| anyhow::anyhow!("无法读取图片 {}: {}", image_path, e))?;
println!("图片读取成功,字节大小: {}", image_bytes.len());
// 【修改点 1】将字节流解码为统一的 DynamicImage
let img = image::load_from_memory(&image_bytes)
.map_err(|e| anyhow::anyhow!("图片解码失败: {}", e))?;
// 【修改点 2】传入统一的 &DynamicImage 引用
let bboxes = Detector::new(&det).predict(&img)?;
// println!("{:?}", det);
println!("检测到的目标数量: {}", bboxes.len());
if bboxes.is_empty() {
println!("未检测到任何目标。");
} else {
// 如果 save_debug_image 报错,记得去把它的入参类型和内部访问也改为 DetectionResult
save_debug_image(
&img,
&bboxes,
"D:/CNWei/CNW/Rust/ddddocr-rs/samples/result.jpg",
)?;
for (i, bbox) in bboxes.iter().enumerate() {
// 【修改点 3】将原来的 bbox[0].. 索引访问改为结构体字段访问
println!("目标 [{}]: {}", i, bbox);
}
}
Ok(())
}
#[test]
fn test_real_slide_match() {
let engine = Slider::new();
// 1. 加载你准备好的测试图
// 假设图片放在项目根目录下的 assets 文件夹
let target_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/hua.png")
.expect("请确保 samples/hua.png 存在");
let bg_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/huatu.png")
.expect("请确保 samples/huatu.png 存在");
// 2. 执行匹配
// 如果是那种带有明显阴影边缘的复杂滑块,建议 simple_target 传 false
let start = std::time::Instant::now();
let result = engine
.slide_match(&target_img, &bg_img, false)
.expect("Slide match 执行失败");
let duration = start.elapsed();
// 3. 打印结果
println!("-------------------------------------------");
println!("{}", result);
println!("耗时: {:?}", duration);
println!("-------------------------------------------");
// 验证基本逻辑:坐标不应为 0 (除非匹配失败)
assert_eq!(result.target_x, 237);
assert_eq!(result.target_y, 77);
assert!(result.confidence > 0.0);
}
#[test]
fn test_real_slide_comparison() {
let engine = Slider::new();
// 1. 加载你准备好的测试图
// 假设图片放在项目根目录下的 assets 文件夹
let target_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/ken.jpg")
.expect("请确保 samples/ken.jpg 存在");
let bg_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/kenyuan.jpg")
.expect("请确保 samples/kenyuan.jpg 存在");
// 2. 执行匹配
// 如果是那种带有明显阴影边缘的复杂滑块,建议 simple_target 传 false
let start = std::time::Instant::now();
let result = engine
.slide_comparison(&target_img, &bg_img)
.expect("Slide match 执行失败");
let duration = start.elapsed();
// 3. 打印结果
println!("-------------------------------------------");
println!("滑块匹配测试结果:");
println!("检测坐标: [x: {}, y: {}]", result.target_x, result.target_y);
println!("置信度: {:.4}", result.confidence);
println!("耗时: {:?}", duration);
println!("-------------------------------------------");
// 验证基本逻辑:坐标不应为 0 (除非匹配失败)
assert_eq!(result.target_x, 171);
assert_eq!(result.target_y, 90);
assert!(result.confidence > 0.0);
}
#[test]
fn test_resolve_shape_logic_direct() {
// 创建一个哑 ModelLoader 实例session 用不上,因为我们直接测私有方法)
let loader = OrtModelLoader::default()
.build_for_path(
// "D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx",
"D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_huashi666_i64.onnx",
)
.expect("建立测试模型图失败");
}

26
ddddocr-tract/Cargo.toml Normal file
View File

@@ -0,0 +1,26 @@
[package]
name = "ddddocr-tract"
version = { workspace = true }
edition = { workspace = true }
license = { workspace = true }
[dependencies]
ddddocr-core = { path = "../ddddocr-core" } # 引入兄弟库
tract-onnx = { workspace = true }
tract-linalg = {workspace = true }
ndarray = { workspace = true } # 继承自工作空间
image = { workspace = true }
serde = { workspace = true }
serde_json = { workspace = true }
anyhow = { workspace = true }
thiserror = { workspace = true } # 刚好可以开始接入你需要的标准库错误处理
[features]
default = []
embed-models = [] # 这是一个留给有特殊需求、且自己下载了模型放入 models/ 目录的人的后门

1
ddddocr-tract/src/det.rs Normal file
View File

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

View File

@@ -0,0 +1,52 @@
use crate::types::Session;
use ddddocr_core::error::{Result, TensorError};
use ddddocr_core::{ DetOutput};
use ddddocr_core::traits::{DetEngine, InferenceEngine};
use ndarray::Ix3;
use tract_onnx::prelude::*;
#[derive(Debug)]
pub struct DetRuntime {
pub session: Session,
}
impl DetRuntime {
pub fn new(session: Session) -> Self {
Self { session }
}
}
impl InferenceEngine for DetRuntime {
type Output = DetOutput; // 明确绑定 OCR 小枚举
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
// let result = self.ocr.run(tvec!(tensor.into()))?;
let tensor = Tensor::from(input_array);
let mut result = self
.session
.run(tvec!(tensor.into()))
.map_err(|_| TensorError::Engine("执行模型推理失败".to_string()))?;
println!("模型输出原始数据: {:?}", result);
// Ok(result.swap_remove(0).into_tensor())
let raw_tensor = result.swap_remove(0).into_tensor();
// raw_tensor.into_plain_array()?
let array_d = raw_tensor.into_plain_array::<f32>().map_err(|_| {
TensorError::Engine("Tract 实体张量无法转换为 ndarray::ArrayD".to_string())
})?;
// 提前利用克隆(Clone)备份好当前未转维度前的真实 shape (Vec<usize>)
let actual_shape = array_d.shape().to_vec();
let array3 = array_d.into_dimensionality::<Ix3>().map_err(|_| {
TensorError::DimensionMismatch {
expected: "3D 检测矩阵 [Batch, Box_Count, Box_Attributes]".to_string(),
actual: actual_shape, // 优雅降维失败时动态捕获
}
})?;
Ok(DetOutput::Detection(array3))
// 在引擎内部消化掉 DatumType 强耦合
}
}
impl DetEngine for DetRuntime {}

View File

@@ -0,0 +1,8 @@
use thiserror::Error;
#[derive(Error, Debug)]
pub enum TensorError {
/// 替换原有的 anyhow::Error明确将 Tract/ONNX 引擎底层报错序列化为干净的 String
#[error("推理引擎内部发生异常: {0}")]
Engine(String),
}

11
ddddocr-tract/src/lib.rs Normal file
View File

@@ -0,0 +1,11 @@
mod det;
mod error;
pub mod loader;
mod ocr;
mod types;
pub use ddddocr_core::{
DetectionResult, Detector, ModelMetadata, Normalization, Ocr, OcrBuilder, SlideResult, Slider,
};
pub use det::session::DetRuntime;
pub use ocr::session::OcrRuntime;

View File

@@ -0,0 +1,7 @@
mod error;
mod metadata;
mod model;
pub use error::{Error, ParseError, Result};
pub use metadata::{ModelMetadataDto, NormalizationDto, Metadata};
pub use model::ModelLoader;

View File

@@ -0,0 +1,50 @@
use tract_onnx::prelude::TractError;
pub type Result<T> = std::result::Result<T, Error>;
#[derive(thiserror::Error, Debug)]
pub enum Error {
/// 解析 ONNX 模型/路径失败(如文件损坏、算子不支持、路径非法)
#[error("解析 ONNX 模型结构失败: {0}")]
ModelParse(#[from] ParseError),
/// 模型计算图优化失败(如常量折叠、形状推导失败)
#[error("优化 Tract 模型图失败: {0}")]
OptimizationFailed(#[source] TractError),
/// 构建可执行 Session 失败(如输入输出 Tensor 类型/形状未确定)
#[error("构建可运行 Tract 实例失败: {0}")]
RunnableBuildFailed(#[source] TractError),
/// JSON 反序列化失败(自动透传 serde_json 报错)
#[error("模型 Metadata JSON 解析失败: {0}")]
JsonParse(#[from] serde_json::Error),
/// 字节流非合法 UTF-8 编码(自动透传 Utf8Error
#[error("Metadata 字节流不是合法的 UTF-8 编码: {0}")]
InvalidUtf8(#[from] std::str::Utf8Error),
#[error("模型元数据解析失败: {0}")]
MetadataParse(String),
/// 承载任何第三方扩展、解密、特定预处理插件在执行时产生的自定义错误
#[error("{0}: {1}")]
Other(String, #[source] Box<dyn std::error::Error + Send + Sync>),
}
impl Error {
/// 方便将任何第三方 Error 包装为 Error::Other
pub fn new<E>(msg: impl Into<String>, err: E) -> Self
where
E: Into<Box<dyn std::error::Error + Send + Sync>>,
{
Self::Other(msg.into(), err.into())
}
}
#[derive(thiserror::Error,Debug)]
pub enum ParseError{
/// 策略 A从文件路径加载失败附带路径上下文信息方便排查是找不到文件还是格式不对
#[error("从路径 '{0}' 加载 ONNX 模型失败: {1}")]
Path(String, #[source] TractError),
/// 策略 B从内存字节流加载失败如 include_bytes! 传入的字节流损坏)
#[error("从内存字节流解析 ONNX 模型失败: {0}")]
Bytes(#[source] TractError),
}

View File

@@ -0,0 +1,93 @@
use crate::loader::error::{Error, Result};
use ddddocr_core::ModelMetadata;
use ddddocr_core::Resize;
use ddddocr_core::{Charset, Normalization};
use serde::Deserialize;
use std::borrow::Cow;
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")] // 支持 json 中写 "zero_to_one" 或 "minus_one_to_one"
pub enum NormalizationDto {
/// 映射到 [0.0, 1.0] -> pixel / 255.0
ZeroToOne,
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
MinusOneToOne,
}
impl From<NormalizationDto> for Normalization {
fn from(dto: NormalizationDto) -> Self {
match dto {
NormalizationDto::ZeroToOne => Normalization::ZeroToOne,
NormalizationDto::MinusOneToOne => Normalization::MinusOneToOne,
}
}
}
/// 仅用于反序列化 JSON 的中间临时结构体DTO
#[derive(Deserialize)]
pub struct ModelMetadataDto {
charset: Vec<String>,
word: bool,
#[serde(alias = "image")]
resize: Vec<i32>,
channel: u8,
/// 新增:允许在配置文件中指定归一化策略。
/// 使用 serde(default) 可以在不配置时提供一个默认值(比如默认 ZeroToOne
#[serde(default = "default_normalization")]
normalization: NormalizationDto,
}
fn default_normalization() -> NormalizationDto {
NormalizationDto::ZeroToOne
}
/// Tract 专属扩展trait 或 工具函数
pub trait Metadata: Sized {
fn from_json_str(json_str: &str) -> Result<Self>;
/// 机制 2从内存字节流加载极大地方便 include_bytes! 或网络下载)
fn from_json_bytes(bytes: &[u8]) -> Result<Self> {
let json_str = std::str::from_utf8(bytes)?;
Self::from_json_str(json_str)
}
}
impl Metadata for ModelMetadata {
// --- 优雅的工厂模式构造器 ---
fn from_json_str(json_str: &str) -> Result<ModelMetadata> {
let dto: ModelMetadataDto = serde_json::from_str(json_str)?;
// 1. 将 DTO 的字符串数组转化为强类型的 Charset
let tokens: Vec<Cow<'static, str>> =
dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
let charset = Charset::new(tokens);
// 2. 解析 resize 策略(重现 Python 的复杂条件判断)
if dto.resize.len() != 2 {
return Err(Error::MetadataParse(
"'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]".to_string(),
));
}
let r0 = dto.resize[0];
let r1 = dto.resize[1];
let resize = if r0 == -1 {
if dto.word {
// 如果 word 为 true且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
Resize::Square(r1 as u32)
} else {
// 如果 word 为 false且包含 -1Python 里是高度固定为 r1宽度按原图比例缩放
Resize::DynamicWidth(r1 as u32)
}
} else {
// 正常的固定宽高
Resize::Fixed(r0 as u32, r1 as u32)
};
Ok(ModelMetadata::new(
charset,
dto.word,
resize,
dto.channel,
dto.normalization.into(),
))
}
}

View File

@@ -0,0 +1,127 @@
use crate::loader::error::{Error, ParseError, Result};
use crate::types::Session;
use ddddocr_core::traits::Loader;
use std::io::Cursor;
use tract_linalg::multithread::{Executor, set_default_executor};
use tract_onnx::onnx;
use tract_onnx::prelude::*;
/// Tract 专用的链式构建器
#[derive(Debug, Clone, Default)]
pub struct ModelLoader {
num_threads: Option<usize>,
}
impl ModelLoader {
/// 可选扩展:设置 CPU 线程数(不提供任何 GPU 相关的 API
pub fn num_threads(mut self, threads: usize) -> Self {
self.num_threads = Some(threads);
self
}
fn setup_tract_threads(&self) {
// 💡 1. 如果设置了线程数,可以通过 tract 的 multithread 配置应用给 model
if let Some(threads) = self.num_threads {
// 在 Tract 中可以通过 set_num_threads 或设置底层环境控制并发
// Tract 0.20+ 版本支持全局/局部线程控制)
let executor = if threads <= 1 {
Executor::SingleThread
} else {
Executor::multithread(threads)
};
set_default_executor(executor);
}
}
}
impl Loader for ModelLoader {
type Session = Session;
type Error = Error;
fn build_for_path<P>(&self, model_path: P) -> Result<Session>
where
P: AsRef<std::path::Path>,
{
self.setup_tract_threads();
let path_ref = model_path.as_ref();
let session = onnx()
.model_for_path(path_ref)
.map_err(|e| ParseError::Path(path_ref.display().to_string(), e))?
// .with_context(|| "加载 ONNX 模型失败,请检查路径是否正确")?
.into_optimized()
.map_err(Error::OptimizationFailed)?
// .with_context(|| "优化 Tract 模型图失败")?
.into_runnable()
.map_err(Error::RunnableBuildFailed)?;
// .with_context(|| "构建可运行 Tract 实例失败")?;
Ok(session)
}
/// 策略 B从内存字节流加载模型配合 include_bytes! 使用)
fn build_from_bytes(&self, model_bytes: &[u8]) -> Result<Session> {
self.setup_tract_threads();
// 使用 std::io::Cursor 将 &[u8] 包装为可读的流(实现 std::io::Read
let mut cursor = Cursor::new(model_bytes);
let session = onnx()
.model_for_read(&mut cursor)
.map_err(ParseError::Bytes)?
// .with_context(|| "从内存字节流解析 ONNX 模型失败")?
.into_optimized()
.map_err(Error::OptimizationFailed)?
// .with_context(|| "优化 Tract 模型图失败")?
.into_runnable()
.map_err(Error::RunnableBuildFailed)?;
// .with_context(|| "构建可运行 Tract 实例失败")?;
Ok(session)
}
}
#[cfg(test)]
mod tests {
use super::*;
/// 辅助函数:动态构建一个简单的 ONNX/Tract 内存模型图用于测试
fn create_test_model() -> std::result::Result<Session, anyhow::Error> {
let mut rect = TypedModel::default();
// 0.21.10 最稳妥的静态 Fact 构建
let input_fact = TypedFact::dt_shape(DatumType::F32, &[1, 3, 224, 224]);
let input_node = rect
.add_source("input", input_fact)
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
rect.set_input_outlets(&[input_node.into()])
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
let typed = rect
.into_optimized()
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
let runnable = typed
.into_runnable()
.map_err(|e| anyhow::anyhow!("{:?}", e))?;
Ok(runnable)
}
// #[test]
// fn test_input_output_shapes_and_type() {
// let session = create_test_model().expect("建立测试模型图失败");
//
// println!("{:?}", ModelLoader::model_info(&session).unwrap());
// // 1. 测试输入维度解析
// }
//
// #[test]
// fn test_resolve_shape_logic_direct() {
// // 创建一个哑 ModelLoader 实例session 用不上,因为我们直接测私有方法)
// let session = create_test_model().expect("建立测试模型图失败");
//
// let dims: Vec<TDim> = vec![TDim::from(1), TDim::from(3), TDim::from(224)];
// // 方案二的精髓:我们直接利用已导出的 ShapeFact 来纯手工验证边界逻辑!
// // 1. 验证纯静态维度是否被正确还原
// let static_shape = ShapeFact::from_dims(dims);
//
// let res = ModelLoader
// ::resolve_shape(&static_shape);
// }
}

1
ddddocr-tract/src/ocr.rs Normal file
View File

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

View File

@@ -0,0 +1,153 @@
use crate::types::Session;
use ddddocr_core::ModelMetadata;
use ddddocr_core::OcrOutput;
use ddddocr_core::error::{DdddError, Result, TensorError};
use ddddocr_core::traits::{InferenceEngine, Info, OcrEngine};
use ddddocr_core::types::{AxisDim, ModelInfo, TensorInfo, TensorType};
use ddddocr_core::utils::normalize_ocr_logits;
use tract_onnx::prelude::{DatumType, OutletId, ShapeFact, TypedModel};
use tract_onnx::prelude::{IntoTensor, Tensor, tvec};
pub struct OcrRuntime {
pub session: Session,
pub metadata: ModelMetadata,
}
impl OcrRuntime {
pub fn new(session: Session, metadata: ModelMetadata) -> Self {
Self { session, metadata }
}
/// 获取模型输入的节点信息列表
/// 提取出来的公共转换逻辑:将一组 OutletId 解析为 TensorInfo 列表
fn resolve_tensors(&self, model: &TypedModel, outlets: &[OutletId]) -> Result<Vec<TensorInfo>> {
outlets
.iter()
.map(|&outlet_id| {
let fact = model.outlet_fact(outlet_id).map_err(DdddError::new)?;
// .map_err(|e| {
// DdddError::InternalError(format!("解析节点 Fact 失败: {:?}", e))
// })?;
let shape = self.resolve_shape(&fact.shape);
let node_name = model.node(outlet_id.node).name.clone();
let tensor_type = match fact.datum_type {
DatumType::F32 => TensorType::F32,
DatumType::I64 => TensorType::I64,
_ => TensorType::Other,
};
Ok(TensorInfo {
name: node_name,
shape,
tensor_type,
})
})
.collect() // 函数式声明:自动传播第一处发生的错误
}
/// 安全还原 Tract 维度至 Vec<AxisDim>
fn resolve_shape(&self, shape_fact: &ShapeFact) -> Vec<AxisDim> {
let tract_shape = shape_fact.to_tvec();
let resolved = tract_shape
.iter()
.map(|dim| {
// 防御性编程:必须同时满足能够转换为 i64 且 大于等于 0
if let Ok(size) = dim.to_i64() {
if size >= 0 {
AxisDim::Static(size as usize)
} else {
// 如果 ONNX 导出时某些动态维度被标记为了 -1安全地作为动态符号捕获
AxisDim::Dynamic(dim.to_string())
}
} else {
AxisDim::Dynamic(dim.to_string())
}
})
.collect();
resolved
}
}
impl OcrEngine for OcrRuntime {
fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
}
impl InferenceEngine for OcrRuntime {
type Output = OcrOutput;
/// 对应 Python 的 _inference
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output, TensorError> {
// tract 的 run 会返回一个 Vec<TValue>,我们通常只需要第一个输出
// let result = self.ocr.run(tvec!(tensor.into()))?;
let tensor = Tensor::from(input_array);
let mut result = self
.session
.run(tvec!(tensor.into()))
.map_err(|_| TensorError::Engine("执行模型推理失败".to_string()))?;
// .context("执行模型推理失败")?;
println!("模型输出原始数据: {:?}", result);
// Ok(result.swap_remove(0).into_tensor())
let raw_tensor = result.swap_remove(0).into_tensor();
// 在引擎内部消化掉 DatumType 强耦合
match raw_tensor.datum_type() {
DatumType::I64 => {
let array_d = raw_tensor
.into_plain_array::<i64>()
.map_err(|_| TensorError::Engine("Tract 无法获取 i64 内存视图".to_string()))?;
// .context("Tract 无法获取 i64 内存视图")?;
// 🌟 提前提取真实维度
let actual_shape = array_d.shape().to_vec();
// 转成标准的 Array1 传给 core
let array1 = array_d
.to_owned()
.into_dimensionality::<ndarray::Ix1>()
.map_err(|_| TensorError::DimensionMismatch {
expected: "1D 字符索引静态矩阵".to_string(),
actual: actual_shape,
})?;
Ok(OcrOutput::Indices(array1))
}
DatumType::F32 => {
let shape = raw_tensor.shape();
println!("模型输出shape数据: {:?}", shape);
// raw_tensor.to_plain_array_view()
let view = raw_tensor
.to_plain_array_view::<f32>()
.map_err(|_| TensorError::Engine("Tract 无法获取 f32 内存视图".to_string()))?;
// 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗
normalize_ocr_logits(view, shape)
}
_ => Err(
// anyhow::anyhow!("不支持的模型输出数据类型: {:?}",raw_tensor.datum_type())
TensorError::UnknownOutputFormat,
),
}
}
}
impl Info for OcrRuntime {
fn input_info(&self) -> Result<Vec<TensorInfo>> {
let model = self.session.model();
let outlets = model.input_outlets().map_err(DdddError::new)?;
self.resolve_tensors(model, outlets)
}
/// 获取模型输出的节点信息列表
fn output_info(&self) -> Result<Vec<TensorInfo>> {
let model = self.session.model();
let outlets = model.output_outlets().map_err(DdddError::new)?;
self.resolve_tensors(model, outlets)
}
/// 获取模型详细元数据信息(对标 Python ddddocr 的 get_model_info
/// 完美包容 [1, 1, 64, image_width] 这样的变长图像模型
/// 获取模型详细元数据信息(代码更紧凑、优雅)
fn model_info(&self) -> Result<ModelInfo> {
Ok(ModelInfo {
inputs: self.input_info()?,
outputs: self.output_info()?,
providers: None,
})
}
}

View File

@@ -0,0 +1,4 @@
use std::sync::Arc;
use tract_onnx::prelude::TypedRunnableModel;
pub type Session = Arc<TypedRunnableModel>;

View File

@@ -0,0 +1,600 @@
use std::borrow::Cow;
use std::fs::File;
use std::path::Path;
use anyhow::anyhow;
use ddddocr_core::Charset;
use ddddocr_core::{Normalization, Resize};
pub const CHARSET_BETA: &[&str] = &[
"", "", "", "", "", "", "", "", "", "", "", "", "", "6", "", "",
"", "", "", "", "", "", "", "", "", "", "", "鴿", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "f", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "²", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "!", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "à", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "鹿", "", "", "", "",
"", "p", "L", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "=", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "Y", "", "", "", "", "", "",
"", "", "w", "", "", "3", "", "F", "", "", "", "", "", "", "", "",
"m", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "耀", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "Θ", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "湿",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "X", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "绿", "", "", "", "",
"", "", "滿", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "G", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "x", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "/", "", "", "", "", "", "", "", "", "", "i", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "椿", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", ",", "", "", "", "", "",
"", "T", "", "", "", "N", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "´", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", " ", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "v", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "c",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "''", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "殿", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "B", "", "", "", "О", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "ɔ", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "\"", "", "", "", "", "", "",
"", "", "", "浿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "n",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", ":",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "#", "", "?", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "Φ", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "Q", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", ";", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "轿", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "H", "", "",
"", "", "", "", "", "", "", "", "", "", "", "趿", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "褿", "", "姿", "", "", "", "", "", "", "", "", "", "", "",
"", "K", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "尿", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "W", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", ">", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "P", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "r", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "%",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "l", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "E", "", "", "", "", "", "蹿", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "И", "", "", "", "Z", "", "",
"", "", "", "寿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "α", "", "", "",
"", "", "", "", "", "", "", "", "", "", "s", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "2", "З", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "Ω", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "@", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "z", "", "", "", "", "", "", "", "", "", "", "", "", "", "访",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "巿", "", "", "", "", "", "D", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "鱿", "", "", "O", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "-", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "西", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "羿",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "麿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "Р", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "ä", "", "", "", "", "广", "", "",
"", "", "", "", "", "", "4", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "忿", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "涿", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "°", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "^", "", "", "", "$", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "槿", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", ")", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "ü", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "仿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "1", "", "", "", "", "", "", "", "", "", "", "", "", "Й",
"", "", "", "", "", "", "", "", "", "", "", "", "亿", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", " ", "", "", "", "", "", "", "",
"", "", "", "", "", "", "t", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "竿",
"", "|", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "β", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "b", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "o", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "Ë",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "è", "", "", "", "", "u", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"÷", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"±", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "9", "", "", "", "", "j", "", "", "0", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "\\", "", "", "", "",
"", "", "", "", "", "", "", "8", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "¥", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "贿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "ò", "", "", "{", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "5", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "岿",
"[", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "驿", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "e", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "A", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "线", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "é", "", "",
"", "", "", "", "", "", "", "~", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "R", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"稿", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "窿", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "g", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "k", "", "", "", "", "", "",
"", "", "", "", "", "鸿", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "退", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "S",
"", "}", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "`", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "怀", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "屿", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "<", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "Я", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "Λ", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "齿", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "+", "", "", "宿", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "I", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "便", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "×", "", "", "",
"", "", "", "", "", "", "", "穿", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "7", "", "", "", "", "",
"", "", "", "", "", "", "", "", ".", "", "d", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "V", "", "", "]", "", "", "", "", "",
"", "", "", "", "", "", "(", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "诿", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "沿", "", "", "", "", "", "", "使", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"·", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "饿", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"J", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "a", "", "", "", "", "", "", "", "", "&", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "h", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "*", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "q", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "_", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "簿", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "罿", "", "П", "",
"U", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "廿", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "馿", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "M", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "y", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "C", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "®", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "", "", "", "", "", "", "", "", "", "", "", "",
"", "", "", "", "婿", "", "", "", "", "", "", "", "", "", "", "",
"", "",
];
pub const CHARSET_OLD: &[&str] = &["", "", "", "", ""];
// pub fn from_builtin_old() -> Self {
// Self::from_static_slice(
// CHARSET_OLD,
// false,
// Resize::DynamicWidth(64),
// 1,
// Normalization::ZeroToOne,
// )
// }
//
// /// 从预设的 Beta 版字符集创建
// pub fn from_builtin_beta() -> Self {
// Self::from_static_slice(
// CHARSET_BETA,
// false,
// Resize::DynamicWidth(64),
// 1,
// Normalization::MinusOneToOne,
// )
// }
// /// 从外部外部 JSON 文件动态加载字符集(在后续优化中移除)
// pub fn from_json_file<P: AsRef<Path>>(path: P) -> anyhow::Result<Self> {
// let path = path.as_ref();
// if !path.exists() {
// return Err(anyhow!("模型元数据配置文件不存在: {:?}", path));
// }
//
// let mut file = File::open(path)?;
// let mut content = String::new();
// file.read_to_string(&mut content)?;
//
// let dto: ModelMetadataDto = serde_json::from_str(&content)
// .map_err(|e| anyhow!("JSON 反序列化失败,请检查字段是否完整: {}", e))?;
//
// // 1. 将 DTO 的字符串数组转化为强类型的 Charset
// let tokens: Vec<Cow<'static, str>> =
// dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
// let charset = Charset::new(tokens);
//
// // 2. 解析 resize 策略(重现 Python 的复杂条件判断)
// if dto.resize.len() != 2 {
// return Err(anyhow!(
// "'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]"
// ));
// }
// let r0 = dto.resize[0];
// let r1 = dto.resize[1];
//
// let resize = if r0 == -1 {
// if dto.word {
// // 如果 word 为 true且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
// Resize::Square(r1 as u32)
// } else {
// // 如果 word 为 false且包含 -1Python 里是高度固定为 r1宽度按原图比例缩放
// Resize::DynamicWidth(r1 as u32)
// }
// } else {
// // 正常的固定宽高
// Resize::Fixed(r0 as u32, r1 as u32)
// };
//
// Ok(Self {
// charset,
// word: dto.word,
// resize,
// channel: dto.channel,
// normalization: dto.normalization,
// })
// }

View File

@@ -1,9 +1,18 @@
use ddddocr_rs::models::slide::Slide;
use ddddocr_rs::{DdddOcr, DdddOcrBuilder}; // 假设你的包名是这个
use image::{DynamicImage, Rgb};
use anyhow::Context;
use ddddocr_core::traits::Loader;
use ddddocr_tract::{DetectionResult, Ocr};
use ddddocr_tract::{Detector, ModelMetadata, Normalization, Slider};
// 假设你的包名是这个
use ddddocr_tract::{DetRuntime, OcrRuntime};
use image::{DynamicImage, ImageBuffer, Luma, Rgb};
use std::fs;
use std::path::Path;
use ddddocr_rs::models::det::DetectionResult;
mod char_slice;
use char_slice::CHARSET_BETA;
use ddddocr_core::Resize;
use ddddocr_tract::loader::ModelLoader as TractModelLoader;
fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
// 1. 先将泛型转为具体的 &Path 引用
@@ -17,8 +26,8 @@ fn load_image<P: AsRef<Path>>(path: P) -> anyhow::Result<image::DynamicImage> {
}
/// 将检测结果绘制在图像上并保存
fn save_debug_image(
dynamic_img: &DynamicImage, // 【优化点 1】直接传入解码好的引用拒绝重复解码
bboxes: &[DetectionResult], // 【修改点 1】类型改为自定义结构体切片
dynamic_img: &DynamicImage, // 【优化点 1】直接传入解码好的引用拒绝重复解码
bboxes: &[DetectionResult], // 【修改点 1】类型改为自定义结构体切片
output_path: &str,
) -> anyhow::Result<()> {
// 删除了原本的 let dynamic_img = image::load_from_memory(image_bytes)?;
@@ -60,24 +69,80 @@ fn save_debug_image(
img.save(output_path)?;
Ok(())
}
#[allow(dead_code)]
fn save_rust_result(result: &ImageBuffer<Luma<f32>, Vec<f32>>, filename: &str) {
let (width, height) = result.dimensions();
// 1. 寻找最值进行归一化
let mut max_val = f32::MIN;
let mut min_val = f32::MAX;
for p in result.pixels() {
if p.0[0] > max_val {
max_val = p.0[0];
}
if p.0[0] < min_val {
min_val = p.0[0];
}
}
// 2. 创建 8 位灰度图
let mut out_buf = ImageBuffer::new(width, height);
for y in 0..height {
for x in 0..width {
let val = result.get_pixel(x, y).0[0];
let normalized = if max_val > min_val {
((val - min_val) / (max_val - min_val) * 255.0) as u8
} else {
0u8
};
out_buf.put_pixel(x, y, Luma([normalized]));
}
}
// 3. 保存
DynamicImage::ImageLuma8(out_buf).save(filename).unwrap();
println!("Rust 结果热力图已保存至: {}", filename);
}
#[test]
fn test_full_classification() {
// 1. 初始化模型
let ocr = DdddOcrBuilder::new().build().expect("模型加载失败");
let session = TractModelLoader::default()
.build_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx")
.expect("模型加载失败");
let metadata = ModelMetadata::from_static_slice(
CHARSET_BETA,
false,
Resize::DynamicWidth(64),
1,
Normalization::MinusOneToOne,
);
// 1. 初始化模型
let ocr_runtime = OcrRuntime::new(session, metadata);
// 2. 加载测试图片
let img = image::open("samples/code2.png").expect("测试图片不存在");
let img =
image::open("D:/CNWei/CNW/Rust/ddddocr-rs/samples/code2.png").expect("测试图片不存在");
// 3. 执行识别
let result = ocr.classification(&img).expect("识别过程出错");
// let result = Ocr::new(&ocr_runtime)
// .predict(&img)
// .expect("识别过程出错")
// .into_text();
let result = Ocr::builder()
.runner(&ocr_runtime)
.predict(&img)
.expect("识别过程出错")
.into_text();
println!("识别结果: {}", result);
assert!(!result.is_empty());
}
#[test]
fn test_det_load() -> anyhow::Result<()> {
let det = DdddOcrBuilder::new().det().build()?;
let image_path = "samples/det1.png";
let det_model = TractModelLoader::default()
.build_for_path("D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_det.onnx")
.expect("模型加载失败");
let det = DetRuntime::new(det_model);
let image_path = "D:/CNWei/CNW/Rust/ddddocr-rs/samples/det1.png";
let image_bytes =
fs::read(image_path).map_err(|e| anyhow::anyhow!("无法读取图片 {}: {}", image_path, e))?;
@@ -88,22 +153,23 @@ fn test_det_load() -> anyhow::Result<()> {
.map_err(|e| anyhow::anyhow!("图片解码失败: {}", e))?;
// 【修改点 2】传入统一的 &DynamicImage 引用
let bboxes = det.detection(&img)?;
println!(":?{}", det);
let bboxes = Detector::new(&det).predict(&img)?;
// println!("{:?}", det);
println!("检测到的目标数量: {}", bboxes.len());
if bboxes.is_empty() {
println!("未检测到任何目标。");
} else {
// 如果 save_debug_image 报错,记得去把它的入参类型和内部访问也改为 DetectionResult
save_debug_image(&img, &bboxes, "samples/result.jpg")?;
save_debug_image(
&img,
&bboxes,
"D:/CNWei/CNW/Rust/ddddocr-rs/samples/result.jpg",
)?;
for (i, bbox) in bboxes.iter().enumerate() {
// 【修改点 3】将原来的 bbox[0].. 索引访问改为结构体字段访问
println!(
"目标 [{}]: x1={}, y1={}, x2={}, y2={}, 分数={:.4}, 类别ID={}",
i, bbox.x1, bbox.y1, bbox.x2, bbox.y2, bbox.score, bbox.class_id
);
println!("目标 [{}]: {}", i, bbox);
}
}
Ok(())
@@ -111,12 +177,14 @@ fn test_det_load() -> anyhow::Result<()> {
#[test]
fn test_real_slide_match() {
let engine = Slide::new();
let engine = Slider::new();
// 1. 加载你准备好的测试图
// 假设图片放在项目根目录下的 assets 文件夹
let target_img = load_image("samples/hua.png").expect("请确保 samples/hua.png 存在");
let bg_img = load_image("samples/huatu.png").expect("请确保 samples/huatu.png 存在");
let target_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/hua.png")
.expect("请确保 samples/hua.png 存在");
let bg_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/huatu.png")
.expect("请确保 samples/huatu.png 存在");
// 2. 执行匹配
// 如果是那种带有明显阴影边缘的复杂滑块,建议 simple_target 传 false
@@ -128,9 +196,7 @@ fn test_real_slide_match() {
// 3. 打印结果
println!("-------------------------------------------");
println!("滑块匹配测试结果:");
println!("检测坐标: [x: {}, y: {}]", result.target_x, result.target_y);
println!("置信度: {:.4}", result.confidence);
println!("{}", result);
println!("耗时: {:?}", duration);
println!("-------------------------------------------");
@@ -142,12 +208,14 @@ fn test_real_slide_match() {
#[test]
fn test_real_slide_comparison() {
let engine = Slide::new();
let engine = Slider::new();
// 1. 加载你准备好的测试图
// 假设图片放在项目根目录下的 assets 文件夹
let target_img = load_image("samples/ken.jpg").expect("请确保 samples/ken.jpg 存在");
let bg_img = load_image("samples/kenyuan.jpg").expect("请确保 samples/kenyuan.jpg 存在");
let target_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/ken.jpg")
.expect("请确保 samples/ken.jpg 存在");
let bg_img = load_image("D:/CNWei/CNW/Rust/ddddocr-rs/samples/kenyuan.jpg")
.expect("请确保 samples/kenyuan.jpg 存在");
// 2. 执行匹配
// 如果是那种带有明显阴影边缘的复杂滑块,建议 simple_target 传 false
@@ -170,3 +238,14 @@ fn test_real_slide_comparison() {
assert_eq!(result.target_y, 90);
assert!(result.confidence > 0.0);
}
#[test]
fn test_resolve_shape_logic_direct() {
// 创建一个哑 ModelLoader 实例session 用不上,因为我们直接测私有方法)
let loader = TractModelLoader::default()
.build_for_path(
// "D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_sml2h3_f32.onnx",
"D:\\CNWei\\CNW\\Rust\\ddddocr-rs\\models\\common_huashi666_i64.onnx",
)
.expect("建立测试模型图失败");
println!("{:?}", loader.model().inputs);
}

View File

@@ -1,5 +1,5 @@
fn main() {
let ocr = ddddocr_rs::DdddOcrBuilder::new().build().unwrap();
let img = image::open("samples/code3.png").unwrap();
println!("Result: {}", ocr.classification(&img).unwrap());
// let ocr = ddddocr_rs::DdddOcrBuilder::new().build().unwrap();
// let img = image::open("samples/code3.png").unwrap();
// println!("Result: {}", ocr.classification(&img).unwrap());
}

View File

@@ -1,184 +0,0 @@
mod charset;
mod model_metadata;
pub mod models;
pub mod utils;
use anyhow::{Result, anyhow};
use image::DynamicImage;
use std::fmt::{Display, Formatter};
// 关键点:直接使用 tract 重导出的 ndarray
use crate::charset::CharRestrict;
use crate::model_metadata::ModelMetadata;
use crate::models::det::DetectionResult;
use crate::utils::color_filter::{ColorPreset, HsvRange};
use models::det::Det;
use models::loader::ModelSession;
use models::ocr::Ocr;
pub enum ModelSpec {
/// 默认 OCR (使用内置路径)
OcrModel,
DetModel,
/// 自定义 OCR (路径由用户提供)
CustomOcrModel {
path: String,
model_metadata: ModelMetadata,
},
}
impl ModelSpec {
// 将默认路径定义为内部关联常量
const DEFAULT_OCR_PATH: &'static str = "models/common_sml2h3_f32.onnx";
const DEFAULT_DET_PATH: &'static str = "models/common_det.onnx";
}
pub enum Runtime {
Ocr(Ocr),
Det(Det),
}
impl Runtime {
// 统一获取描述的方法
pub fn desc(&self) -> String {
match self {
Runtime::Ocr(s) => s.desc(), // 调用 Ocr 结构体的方法
Runtime::Det(s) => s.desc(), // 调用 Det 结构体的方法
}
}
}
pub struct DdddOcrBuilder {
mode: ModelSpec,
}
impl DdddOcrBuilder {
pub fn new() -> Self {
Self {
mode: ModelSpec::OcrModel,
}
}
/// 切换为检测模式
pub fn det(mut self) -> Self {
self.mode = ModelSpec::DetModel;
self
}
/// 设置自定义 OCR 路径
pub fn custom_ocr(mut self, path: String, model_metadata: ModelMetadata) -> Self {
// 直接重写枚举,替换掉之前的 Ocr 或 Det
self.mode = ModelSpec::CustomOcrModel {
path,
model_metadata,
};
self
}
/// 核心初始化逻辑
pub fn build(self) -> Result<DdddOcr> {
let runtime = match self.mode {
ModelSpec::OcrModel => Runtime::Ocr(Ocr::new(
ModelSpec::DEFAULT_OCR_PATH.into(),
ModelMetadata::from_builtin_beta(),
)?),
ModelSpec::DetModel => Runtime::Det(Det::new(ModelSpec::DEFAULT_DET_PATH.into())?),
ModelSpec::CustomOcrModel {
path,
model_metadata,
} => Runtime::Ocr(Ocr::new(path, model_metadata)?),
};
Ok(DdddOcr { runtime })
}
}
pub struct DdddOcr {
runtime: Runtime,
}
impl Display for DdddOcr {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "DdddOcr(session: {})", self.runtime.desc())
}
}
impl DdddOcr {
pub fn classification(&self, img: &DynamicImage) -> Result<String> {
match &self.runtime {
// Runtime::Ocr(s) => s.predict(img).run(),
// Runtime::Ocr(s) => s.predictor().probability(false).predict(img),
// Runtime::Ocr(s) => {
// let predictor = s.predictor();
// let restricted = predictor.charset_restrict(&CharRestrict::Lowercase);
// let a = restricted.valid_tokens();
// println!("{:?}", a);
// Ok("".to_string())
// }
Runtime::Ocr(s) => {
let res = s.predictor().probability(true).predict(img)?;
println!("{}", res);
Ok(res.to_string())
}
// Runtime::Ocr(s) => s.predictor().charset_restrict(&CharRestrict::Digit).predict(img),
// Runtime::Ocr(s) => s.predictor().color_filter(&ColorPreset::Custom(vec![
// // 错误:下界 (82, 221, 14) 没问题
// // 但上界的 H 通道写成了 240超过了 180 的法定上限!
// HsvRange::new((82, 221, 14), (240, 203, 82)),
// ])).predict(img),
Runtime::Det(_) => Err(anyhow::anyhow!("当前模型是检测模型,无法执行 OCR")),
}
}
pub fn detection(&self, img: &DynamicImage) -> Result<Vec<DetectionResult>> {
match &self.runtime {
Runtime::Det(s) => s.predict(img),
Runtime::Ocr(_) => Err(anyhow::anyhow!("当前模型是 OCR 模型,无法执行检测")),
}
}
}
// struct Classification {}
// #[derive(Debug)]
// struct ClassificationBuilder {
// img: DynamicImage,
// png_fix: bool,
// color_filter_colors: Option<Vec<ColorRange>>,
// color_filter_custom_ranges: Option<Vec<ColorRange>>,
// }
// impl ClassificationBuilder {
// pub fn new(img: DynamicImage) -> Self {
// ClassificationBuilder {
// img,
// png_fix: false,
// color_filter_colors: None,
// color_filter_custom_ranges: None,
// }
// }
// pub fn png_fix(mut self, value: bool) -> Self {
// self.png_fix = value;
// self
// }
// pub fn color_filter_colors(mut self, value: Vec<ColorRange>) -> Self {
// self.color_filter_colors = Some(value);
// self
// }
// pub fn color_filter_custom_ranges(mut self, value: Vec<ColorRange>) -> Self {
// self.color_filter_custom_ranges = Some(value);
// self
// }
// pub fn build(self) -> Classification {
// Classification {}
// }
// }
#[cfg(test)]
mod tests {
#[test]
fn test_ctc_decode_indices() {
// 模拟一个 DdddOcr 实例(如果 decode 不依赖 session可以设为相关函数
// 这里假设你的 decode_ctc 是公开或内部可访问的
let input = vec![1, 1, 0, 1, 2, 2, 0, 2];
// 逻辑:[1, 1] -> 1, [0] -> 跳过, [1] -> 1, [2, 2] -> 2, [0] -> 跳过, [2] -> 2
// 预期结果索引应该是 [1, 1, 2, 2] 对应的字符
// 具体的断言取决于你的 CHARSET_BETA
// let result = dddd.ctc_decode_indices(&input);
// assert_eq!(result, "AABB");
}
}

View File

@@ -1,166 +0,0 @@
use crate::charset::{CHARSET_BETA, CHARSET_OLD, Charset};
use anyhow::{Result, anyhow};
use serde::Deserialize;
use std::borrow::Cow;
use std::collections::{HashMap, HashSet};
use std::fs::File;
use std::io::Read;
use std::path::Path;
// =====================================================================
// 1. 辅助定义的枚举与结构体
// =====================================================================
#[derive(Debug, Clone, Copy, Deserialize)]
#[serde(rename_all = "snake_case")] // 支持 json 中写 "zero_to_one" 或 "minus_one_to_one"
pub enum Normalization {
/// 映射到 [0.0, 1.0] -> pixel / 255.0
ZeroToOne,
/// 映射到 [-1.0, 1.0] -> (pixel / 255.0 - 0.5) / 0.5
MinusOneToOne,
}
impl Normalization {
/// 统一归一化计算逻辑
#[inline(always)]
pub fn normalize(&self, pixel: f32) -> f32 {
match self {
Normalization::ZeroToOne => pixel / 255.0,
Normalization::MinusOneToOne => (pixel / 255.0 - 0.5) / 0.5,
}
}
}
/// 图像缩放策略枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Resize {
/// 固定宽高,例如 (64, 64)
Fixed(u32, u32),
/// 高度固定,宽度根据原始比例动态计算(对应 Python 的 [-1, H]
DynamicWidth(u32),
/// 单字识别的正方形切图(对应 Python 的 word 为 True 且 [-1, H]
Square(u32),
}
/// 仅用于反序列化 JSON 的中间临时结构体DTO
#[derive(Deserialize)]
struct ModelMetadataDto {
charset: Vec<String>,
word: bool,
#[serde(alias = "image")]
resize: Vec<i32>,
channel: u8,
/// 新增:允许在配置文件中指定归一化策略。
/// 使用 serde(default) 可以在不配置时提供一个默认值(比如默认 ZeroToOne
#[serde(default = "default_normalization")]
normalization: Normalization,
}
fn default_normalization() -> Normalization {
Normalization::ZeroToOne
}
#[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 from_builtin_old() -> Self {
Self::from_static_slice(
CHARSET_OLD,
false,
Resize::DynamicWidth(64),
1,
Normalization::ZeroToOne,
)
}
/// 从预设的 Beta 版字符集创建
pub fn from_builtin_beta() -> Self {
Self::from_static_slice(
CHARSET_BETA,
false,
Resize::DynamicWidth(64),
1,
Normalization::MinusOneToOne,
)
}
/// 通用的静态切片转换构造器
pub fn from_static_slice(
slice: &[&'static str],
word: bool,
resize: Resize,
channel: u8,
normalization: Normalization,
) -> Self {
let tokens: Vec<Cow<'static, str>> = slice.iter().map(|&s| Cow::Borrowed(s)).collect();
Self {
charset: Charset::new(tokens),
word,
resize,
channel,
normalization,
}
}
/// 从外部外部 JSON 文件动态加载字符集
pub fn from_json_file<P: AsRef<Path>>(path: P) -> Result<Self> {
let path = path.as_ref();
if !path.exists() {
return Err(anyhow!("模型元数据配置文件不存在: {:?}", path));
}
let mut file = File::open(path)?;
let mut content = String::new();
file.read_to_string(&mut content)?;
let dto: ModelMetadataDto = serde_json::from_str(&content)
.map_err(|e| anyhow!("JSON 反序列化失败,请检查字段是否完整: {}", e))?;
// 1. 将 DTO 的字符串数组转化为强类型的 Charset
let tokens: Vec<Cow<'static, str>> =
dto.charset.into_iter().map(|s| Cow::Owned(s)).collect();
let charset = Charset::new(tokens);
// 2. 解析 resize 策略(重现 Python 的复杂条件判断)
if dto.resize.len() != 2 {
return Err(anyhow!(
"'resize (or image)' 字段必须是包含两个元素的数组,例如 [-1, 64]"
));
}
let r0 = dto.resize[0];
let r1 = dto.resize[1];
let resize = if r0 == -1 {
if dto.word {
// 如果 word 为 true且包含 -1Python 里是 resize 为 (r1, r1) 的正方形
Resize::Square(r1 as u32)
} else {
// 如果 word 为 false且包含 -1Python 里是高度固定为 r1宽度按原图比例缩放
Resize::DynamicWidth(r1 as u32)
}
} else {
// 正常的固定宽高
Resize::Fixed(r0 as u32, r1 as u32)
};
Ok(Self {
charset,
word: dto.word,
resize,
channel: dto.channel,
normalization: dto.normalization,
})
}
}

View File

@@ -1,40 +0,0 @@
pub trait ModelArgs {
// 获取模型路径
fn model_path(&self) -> &str;
// 获取字符集(由于 Det 没有,所以返回 Option
fn charset(&self) -> Option<&str>;
}
pub struct HasCharset {
pub charset: String,
} // 给 Ocr 和 Custom 用
pub struct NoCharset; // 给 Det 用
pub struct Model<T> {
pub path: String,
pub metadata: T,
}
// 针对有字符集的模型 (Ocr / Custom)
impl ModelArgs for Model<HasCharset> {
fn model_path(&self) -> &str {
&self.path
}
fn charset(&self) -> Option<&str> {
Some(&self.metadata.charset)
}
}
// 针对没有字符集的模型 (Det)
impl ModelArgs for Model<NoCharset> {
fn model_path(&self) -> &str {
&self.path
}
fn charset(&self) -> Option<&str> {
None
}
}
pub type OcrModel = Model<HasCharset>;
pub type DetModel = Model<NoCharset>;
pub type CustomModel = Model<HasCharset>; // Ocr 和 Custom 逻辑一致,可以复用

View File

@@ -1,40 +0,0 @@
use anyhow::Context;
use image::DynamicImage;
use tract_onnx::onnx;
use tract_onnx::prelude::*;
// 关键点:直接使用 tract 重导出的 ndarray
use crate::utils::image_io::png_rgba_white_preprocess;
use crate::utils::image_processor::{convert_to_grayscale, resize_image};
use std::collections::HashMap;
use tract_onnx::prelude::tract_ndarray::s;
/// OCR 模型:包含路径和字符集
pub enum ModelType {
Ocr,
Det,
Custom,
}
// 定义统一的 trait
pub trait ModelSession {
fn get_model_type(&self) -> ModelType;
fn desc(&self) -> String;
}
pub struct ModelLoader {
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
}
impl ModelLoader {
pub fn load_model<P>(model_path: P) -> anyhow::Result<Self>
where
P: AsRef<std::path::Path>,
{
let session = onnx()
.model_for_path(model_path)
.with_context(|| "加载 ONNX 模型失败,请检查路径是否正确")?
.into_optimized()?
.into_runnable()?;
Ok(Self { session })
}
}

View File

@@ -1,5 +0,0 @@
pub mod base;
pub mod loader;
pub mod ocr;
pub mod det;
pub mod slide;

View File

@@ -1,264 +0,0 @@
use anyhow::{Context, Result, anyhow, bail};
use base64::{Engine as _, engine::general_purpose};
use image::{DynamicImage, GenericImageView, ImageBuffer, ImageFormat, Luma, Rgb, RgbImage, Rgba};
use std::fs;
use std::path::{Path, PathBuf};
use tract_onnx::prelude::tract_ndarray::{Array3, ArrayD, ArrayViewD};
#[derive(Debug)]
pub enum ColorMode {
RGB,
RGBA,
L,
}
/// 定义支持的输入类型枚举
pub enum ImageInput {
Bytes(Vec<u8>),
Array(ArrayD<u8>), // 对应 numpy 数组
Path(PathBuf),
Base64(String),
DynamicImage(DynamicImage),
}
/// 模拟 Python 的 load_image_from_input
#[allow(dead_code)]
pub fn load_image_from_input(img_input: ImageInput) -> Result<DynamicImage> {
match img_input {
// 2. 处理字节流 (Bytes)
ImageInput::Bytes(bytes) => {
image::load_from_memory(&bytes).context("Failed to load utils from bytes")
}
// 1. 已经是 DynamicImage
ImageInput::DynamicImage(i) => Ok(i),
// 5. 处理 ndarray (Numpy-like)
// 假设输入是 HWC 格式的 Array3<u8>
ImageInput::Array(a) => numpy_to_pil_image(a.view()),
// 4. 处理 Base64 字符串
ImageInput::Base64(b) => base64_to_image(&b),
// 3. 处理文件路径 (Path)
ImageInput::Path(p) => image::open(p).context("Failed to open utils from path"),
}
}
fn base64_to_image(b64_str: &str) -> Result<DynamicImage> {
// 过滤掉可能存在的 base64 前缀,例如 "data:utils/png;base64,"
let clean_b64 = if let Some(pos) = b64_str.find(",") {
&b64_str[pos + 1..]
} else {
&b64_str
};
let bytes = general_purpose::STANDARD
.decode(clean_b64.trim())
.map_err(|e| anyhow!("Base64 decode error: {}", e))?;
image::load_from_memory(&bytes).context("Failed to load utils from decoded base64")
}
/// 读取图片文件并转换为 base64 编码字符串
/// 对应 Python 版 get_img_base64
pub fn get_img_base64<P: AsRef<Path>>(image_path: P) -> Result<String> {
// 1. 读取文件原始字节流
// 使用 AsRef<Path> 泛型可以让函数同时支持 String, &str, PathBuf 等类型
let image_data = fs::read(&image_path)
.with_context(|| format!("Failed to read utils file: {:?}", image_path.as_ref()))?;
// 2. 进行 Base64 编码
// 使用 STANDARD 引擎对齐 Python 的 base64.b64encode
let b64_string = general_purpose::STANDARD.encode(image_data);
Ok(b64_string)
}
/// 封装数组转图像的逻辑,对齐 Python 版 _numpy_to_pil_image
fn numpy_to_pil_image(array: ArrayViewD<u8>) -> Result<DynamicImage> {
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();
match dim {
// 对应 Python: len(array.shape) == 2 (灰度图 H, W)
2 => {
let (h, w) = (shape[0], shape[1]);
ImageBuffer::<Luma<u8>, _>::from_raw(w as u32, h as u32, raw_data)
.map(DynamicImage::ImageLuma8)
.ok_or_else(|| anyhow!("Failed to create Luma utils from 2D array"))
}
// 对应 Python: len(array.shape) == 3 (H, W, C)
3 => {
let (h, w, c) = (shape[0], shape[1], shape[2]);
match c {
// 对应 Python: array.shape[2] == 1 (单通道 H, W, 1)
1 => ImageBuffer::<Luma<u8>, _>::from_raw(w as u32, h as u32, raw_data)
.map(DynamicImage::ImageLuma8),
// 对应 Python: array.shape[2] == 3 (RGB H, W, 3)
3 => ImageBuffer::<Rgb<u8>, _>::from_raw(w as u32, h as u32, raw_data)
.map(DynamicImage::ImageRgb8),
// 对应 Python: array.shape[2] == 4 (RGBA H, W, 4)
4 => ImageBuffer::<Rgba<u8>, _>::from_raw(w as u32, h as u32, raw_data)
.map(DynamicImage::ImageRgba8),
_ => {
return Err(anyhow!("不支持的通道数: {}", c));
}
}
.ok_or_else(|| anyhow!("转换彩色图失败"))
}
_ => Err(anyhow!("不支持的数组维度: {},仅支持 2D 或 3D", dim)),
}
}
/// 对应 Python 的 png_rgba_black_preprocess
/// 将带有透明通道的图片转换为白色背景的 RGB 图片
pub fn png_rgba_white_preprocess(img: &DynamicImage) -> DynamicImage {
// 1. 检查是否包含透明通道,如果没有,直接克隆并返回
if !img.color().has_alpha() {
return DynamicImage::ImageRgb8(img.to_rgb8());
}
let (width, height) = img.dimensions();
// 2. 创建一个新的 RGB 图像缓冲,默认填充为白色 (255, 255, 255)
let mut background = ImageBuffer::from_pixel(width, height, Rgb([255u8, 255u8, 255u8]));
// 3. 获取原图的 RGBA 视图
let rgba_img = img.to_rgba8();
// 4. 遍历像素并手动进行 Alpha 混合
// 对应 Python 的 utils.paste(img, ..., mask=img)
// 使用 enumerate_pixels_mut 同时获取坐标和背景像素的可变引用,减少查找开销
for (x, y, bg_pixel) in background.enumerate_pixels_mut() {
// 安全性说明x, y 源自 background 尺寸,与 rgba_img 一致get_pixel 是安全的
let src_pixel = rgba_img.get_pixel(x, y);
let alpha_u8 = src_pixel[3];
match alpha_u8 {
// 情况 A完全不透明直接覆盖背景色
255 => {
bg_pixel.0 = [src_pixel[0], src_pixel[1], src_pixel[2]];
}
// 情况 B完全透明保持背景色白色无需操作
0 => {
continue;
}
// 情况 C半透明进行 Alpha 混合计算
_ => {
let alpha = alpha_u8 as f32 / 255.0;
let inv_alpha = 1.0 - alpha;
bg_pixel[0] = (src_pixel[0] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
bg_pixel[1] = (src_pixel[1] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
bg_pixel[2] = (src_pixel[2] as f32 * alpha + 255.0 * inv_alpha).round() as u8;
}
}
}
DynamicImage::ImageRgb8(background)
}
pub fn image_to_numpy(image: &DynamicImage, mode: ColorMode) -> Result<Array3<u8>> {
// 1. 模式转换 (对应 utils.convert(target_mode)),此函数在时保留看后续优化是否需要替代image_to_ndarray
// Rust utils 库通过 to_rgb8, to_luma8 等方法实现转换
let (width, height) = image.dimensions();
let (channels, raw) = match mode {
ColorMode::RGB => (3, image.to_rgb8().into_raw()),
ColorMode::L => (1, image.to_luma8().into_raw()),
ColorMode::RGBA => (4, image.to_rgba8().into_raw()),
};
Array3::from_shape_vec((height as usize, width as usize, channels), raw)
.map_err(|e| anyhow!("Failed to build ndarray: {}", e))
}
pub fn numpy_to_image(array: ArrayViewD<u8>, mode: ColorMode) -> Result<DynamicImage> {
let shape = array.shape();
// 1. 基础维度检查 (必须是 H, W, C 三维数组)
if shape.len() != 3 {
bail!("Expected a 3D array (H, W, C), but got {}D", shape.len());
}
let height = shape[0] as u32;
let width = shape[1] as u32;
let channels = shape[2];
// 2. 检查通道数是否与模式匹配
let expected_channels = match mode {
ColorMode::L => 1,
ColorMode::RGB => 3,
ColorMode::RGBA => 4,
};
if channels != expected_channels {
bail!(
"Mode {:?} expects {} channels, but array has {}",
mode,
expected_channels,
channels
);
}
// 确保数据连续性 (C-order)
let standard = array.as_standard_layout();
let (raw_data, _) = standard.to_owned().into_raw_vec_and_offset();
match mode {
ColorMode::L => ImageBuffer::<Luma<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageLuma8),
ColorMode::RGB => ImageBuffer::<Rgb<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageRgb8),
ColorMode::RGBA => ImageBuffer::<Rgba<u8>, _>::from_raw(width, height, raw_data)
.map(DynamicImage::ImageRgba8),
}
.ok_or_else(|| anyhow!("Failed to construct ImageBuffer. Buffer size might be incorrect."))
}
pub fn image_to_ndarray(img: &DynamicImage) -> Array3<u8> {
let (width, height) = img.dimensions();
// 1. 强制转为 RGB8 (丢弃 Alpha 通道,与 Python 的 target_mode='RGB' 对齐)
let rgb_img = img.to_rgb8();
// 2. 获取原始像素数据
let raw_data = rgb_img.into_raw();
// 3. 构造数组 (通道数改为 3)
Array3::from_shape_vec((height as usize, width as usize, 3), raw_data)
.expect("Failed to construct ndarray from utils") // 建议显式报错,而不是返回全黑图
}
#[allow(dead_code)]
fn save_rust_result(result: &ImageBuffer<Luma<f32>, Vec<f32>>, filename: &str) {
let (width, height) = result.dimensions();
// 1. 寻找最值进行归一化
let mut max_val = f32::MIN;
let mut min_val = f32::MAX;
for p in result.pixels() {
if p.0[0] > max_val {
max_val = p.0[0];
}
if p.0[0] < min_val {
min_val = p.0[0];
}
}
// 2. 创建 8 位灰度图
let mut out_buf = ImageBuffer::new(width, height);
for y in 0..height {
for x in 0..width {
let val = result.get_pixel(x, y).0[0];
let normalized = if max_val > min_val {
((val - min_val) / (max_val - min_val) * 255.0) as u8
} else {
0u8
};
out_buf.put_pixel(x, y, Luma([normalized]));
}
}
// 3. 保存
DynamicImage::ImageLuma8(out_buf).save(filename).unwrap();
println!("Rust 结果热力图已保存至: {}", filename);
}

View File

@@ -1,37 +0,0 @@
use image::{DynamicImage, GrayImage, imageops::FilterType};
use anyhow::Result;
/// 对应 Python 的 convert_to_grayscale
/// 将图像转换为灰度图 (L模式)
pub fn convert_to_grayscale(image: &DynamicImage) -> GrayImage {
// Rust utils 库的 to_luma8 会根据标准的亮度公式进行转换
image.to_luma8()
}
/// 对应 Python 的 resize_image
/// 调整图像尺寸。当前版本仅实现 keep_aspect_ratio=false
pub fn resize_image(
image: &DynamicImage,
target_width: u32,
target_height: u32,
// resample 参数我们直接使用 FilterTypeLanczos3 是最接近 Python LANCZOS 的
) -> DynamicImage {
// image::imageops::resize 的最高层封装
// FilterType::Lanczos3 与 Python Pillow 的 Image.LANCZOS 算法完全对齐,缩放质量最高
image.resize_exact(target_width, target_height, FilterType::Lanczos3)
}
// pub fn resize_image(
// image: &GrayImage,
// target_width: u32,
// target_height: u32,
// // resample 参数我们直接使用 FilterTypeLanczos3 是最接近 Python LANCZOS 的
// ) -> GrayImage {
// // 使用 resize 算法进行精确缩放
// image::imageops::resize(
// image,
// target_width,
// target_height,
// FilterType::Lanczos3
// )
// }

View File

@@ -1,4 +0,0 @@
pub mod image_io;
pub mod image_processor;
pub mod cv_ops;
pub mod color_filter;