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, Graph>>, pub model_metadata: ModelMetadata, } impl OcrSession { pub fn new

(model_path: P, model_metadata: ModelMetadata) -> Result where P: AsRef, { 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 { 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) -> Result { // tract 的 run 会返回一个 Vec,我们通常只需要第一个输出 // 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::() .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::() .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::() .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) ), } } }