message / app /utils /ocr_engine.py
hunian
Refactor video audio extractor plugin and enhance functionality
e5e756a
Raw
History Blame Contribute Delete
4.82 kB
"""
OCR 引擎模块 - 封装 PaddleOCR 单例
支持 PP-OCRv6 medium det/rec 成对配置,从 HuggingFace 加载权重。
"""
import logging
import os
from typing import Optional, Dict, Any
from pathlib import Path
from app.plugins.model_probe import probe_provider, ProviderAvailability
logger = logging.getLogger(__name__)
# 全局 OCR 引擎实例
_ocr_engine: Optional[object] = None
_ocr_model_info: Optional[Dict[str, Any]] = None
# PP-OCRv6 模型配置
PP_OCRV6_MODELS = {
"det": {
"model_name": "PP-OCRv6_medium_det",
"hf_model_id": "PaddlePaddle/PP-OCRv6_medium_det_safetensors",
"provider": "ppocrv6",
},
"rec": {
"model_name": "PP-OCRv6_medium_rec",
"hf_model_id": "PaddlePaddle/PP-OCRv6_medium_rec_safetensors",
"provider": "ppocrv6_rec",
},
}
def get_ocr_model_config() -> Dict[str, Any]:
"""获取 OCR 模型配置
从环境变量或默认配置读取 det/rec 模型 ID。
Returns:
Dict 包含 det 和 rec 的模型配置
"""
det_model = os.getenv("OCR_DET_MODEL", PP_OCRV6_MODELS["det"]["model_name"])
rec_model = os.getenv("OCR_REC_MODEL", PP_OCRV6_MODELS["rec"]["model_name"])
return {
"det": {
"model_name": det_model,
"hf_model_id": PP_OCRV6_MODELS["det"]["hf_model_id"],
"provider": PP_OCRV6_MODELS["det"]["provider"],
},
"rec": {
"model_name": rec_model,
"hf_model_id": PP_OCRV6_MODELS["rec"]["hf_model_id"],
"provider": PP_OCRV6_MODELS["rec"]["provider"],
},
}
def probe_ocr_providers() -> Dict[str, ProviderAvailability]:
"""探测 OCR provider 的可用性
检查 paddleocr 和 paddle 依赖是否可用。
Returns:
Dict 包含 det 和 rec provider 的可用性状态
"""
config = get_ocr_model_config()
results = {}
for key, cfg in config.items():
status = probe_provider(
provider=cfg["provider"],
model_id=cfg["hf_model_id"],
check_dependencies=["paddleocr", "paddle"],
)
results[key] = status.status
return results
def get_ocr_engine():
"""
获取全局 OCR 引擎实例(单例模式)
使用 PP-OCRv6 medium det/rec 成对配置,从 HuggingFace 加载权重。
Returns:
PaddleOCR 实例
"""
global _ocr_engine, _ocr_model_info
if _ocr_engine is None:
config = get_ocr_model_config()
det_cfg = config["det"]
rec_cfg = config["rec"]
logger.info(f"正在初始化 OCR 引擎: det={det_cfg['model_name']}, rec={rec_cfg['model_name']}")
try:
from paddleocr import PaddleOCR
# 使用 PP-OCRv6 medium det/rec 成对配置
# 使用 transformers engine 避免 paddle inference 的兼容性问题
_ocr_engine = PaddleOCR(
text_detection_model_name=det_cfg["model_name"],
text_recognition_model_name=rec_cfg["model_name"],
use_doc_orientation_classify=False, # 禁用文档方向分类
use_doc_unwarping=False, # 禁用文档矫正
use_textline_orientation=False, # 禁用文本行方向分类
lang='ch', # 中文模型
engine="transformers", # 使用 transformers engine
)
# 记录模型信息
_ocr_model_info = {
"det_model": det_cfg["model_name"],
"rec_model": rec_cfg["model_name"],
"det_hf_id": det_cfg["hf_model_id"],
"rec_hf_id": rec_cfg["hf_model_id"],
}
logger.info("OCR 引擎初始化完成")
except Exception as e:
logger.error(f"OCR 引擎初始化失败: {e}")
raise
return _ocr_engine
def get_ocr_model_info() -> Dict[str, Any]:
"""获取当前 OCR 模型信息
Returns:
Dict 包含 det_model、rec_model 等信息
"""
global _ocr_model_info
if _ocr_model_info is None:
config = get_ocr_model_config()
_ocr_model_info = {
"det_model": config["det"]["model_name"],
"rec_model": config["rec"]["model_name"],
"det_hf_id": config["det"]["hf_model_id"],
"rec_hf_id": config["rec"]["hf_model_id"],
}
return _ocr_model_info
def preload_ocr_engine():
"""
预加载 OCR 引擎
在应用启动时调用,提前加载模型以加快后续响应速度
"""
try:
engine = get_ocr_engine()
logger.info("OCR 模型预加载成功")
return engine
except Exception as e:
logger.error(f"OCR 模型预加载失败: {e}")
raise