- 新增 Other变体以及构造函数new - 剥离图像预处理中的 Base64 相关错误至业务层处理 - 引入强类型 `LogitsDimensionMismatch` 替代不便匹配的字符串错误 - 优化 `normalize_ocr_logits` 的转换流程,兼顾零拷贝性能与精细化报错 - 优化 全库错误处理
105 lines
4.1 KiB
Rust
105 lines
4.1 KiB
Rust
use crate::loader::ModelLoader;
|
|
use anyhow::Context;
|
|
use ddddocr_core::error::{DdddError, Result, TensorErrorReason};
|
|
use ddddocr_core::utils::normalize_ocr_logits;
|
|
use ddddocr_core::{InferenceEngine, ModelMetadata, OcrEngine, OcrOutput};
|
|
use ndarray::s;
|
|
use std::path::Path;
|
|
use tract_onnx::prelude::DatumType;
|
|
use tract_onnx::prelude::{Graph, IntoTensor, RunnableModel, Tensor, TypedFact, TypedOp, tvec};
|
|
|
|
pub struct OcrSession {
|
|
pub session: RunnableModel<TypedFact, Box<dyn TypedOp>, Graph<TypedFact, Box<dyn TypedOp>>>,
|
|
pub model_metadata: ModelMetadata,
|
|
}
|
|
impl OcrSession {
|
|
pub fn new<P>(model_path: P, model_metadata: ModelMetadata) -> Result<Self>
|
|
where
|
|
P: AsRef<Path>,
|
|
{
|
|
let session = ModelLoader::model_for_path(model_path)?.session;
|
|
Ok(Self {
|
|
session,
|
|
model_metadata,
|
|
})
|
|
}
|
|
|
|
pub fn model_from_bytes(model_bytes: &[u8], model_metadata: ModelMetadata) -> Result<Self> {
|
|
let session = ModelLoader::model_from_bytes(model_bytes)?.session;
|
|
Ok(Self {
|
|
session,
|
|
model_metadata,
|
|
})
|
|
}
|
|
}
|
|
impl OcrEngine for OcrSession {
|
|
fn metadata(&self) -> &ModelMetadata {
|
|
&self.model_metadata
|
|
}
|
|
}
|
|
impl InferenceEngine for OcrSession {
|
|
type Output = OcrOutput;
|
|
/// 对应 Python 的 _inference
|
|
fn inference(&self, input_array: ndarray::Array4<f32>) -> Result<Self::Output> {
|
|
// 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(|_| {
|
|
DdddError::Inference(TensorErrorReason::EngineError(
|
|
"执行模型推理失败".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_array::<i64>()
|
|
.map_err(|_| {
|
|
DdddError::Inference(TensorErrorReason::EngineError(
|
|
"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(|_| {
|
|
DdddError::Inference(TensorErrorReason::TensorDimensionMismatch {
|
|
expected: "1D 字符索引静态矩阵".to_string(),
|
|
actual: actual_shape,
|
|
})
|
|
})?;
|
|
Ok(OcrOutput::Indices(array1))
|
|
}
|
|
DatumType::F32 => {
|
|
let shape = raw_tensor.shape();
|
|
println!("模型输出shape数据: {:?}", shape);
|
|
let view = raw_tensor
|
|
.to_array_view::<f32>()
|
|
.map_err(|_| {
|
|
DdddError::Inference(TensorErrorReason::EngineError(
|
|
"Tract 无法获取 f32 内存视图".to_string(),
|
|
))
|
|
})?;
|
|
// 1. 极其纯粹的、无拷贝的多维 Shape 压扁清洗
|
|
normalize_ocr_logits(view, shape)
|
|
}
|
|
_ => Err(
|
|
// anyhow::anyhow!("不支持的模型输出数据类型: {:?}",raw_tensor.datum_type())
|
|
DdddError::Inference(TensorErrorReason::UnknownOutputFormat)
|
|
),
|
|
}
|
|
}
|
|
}
|