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 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, word: bool, #[serde(alias = "image")] resize: Vec, 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; /// 机制 2:从内存字节流加载(极大地方便 include_bytes! 或网络下载) fn from_json_bytes(bytes: &[u8]) -> Result { 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 { let dto: ModelMetadataDto = serde_json::from_str(json_str)?; // 1. 将 DTO 的字符串数组转化为强类型的 Charset let tokens: Vec> = 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,且包含 -1,Python 里是 resize 为 (r1, r1) 的正方形 Resize::Square(r1 as u32) } else { // 如果 word 为 false,且包含 -1,Python 里是高度固定为 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(), )) } }