scene-detection / module /model_loader.py
jslmmfboom-coder
Revert: 回滚ONNX(HF Spaces上反而更慢), 移除可视化文字叠加
27fd084
Raw
History Blame Contribute Delete
4.04 kB
"""模型加载器 — CPU 部署版(无 MASt3R)
保留 DINOv2 / PaddleOCR / BGE 加载逻辑,
移除 MASt3R 加载(load_model 返回 None)。
"""
import os
import sys
import importlib
import importlib.util
# 屏蔽 tensorflow/jax 导入(transformers 间接依赖,与 numpy 1.x 不兼容)
_original_find_spec = importlib.util.find_spec
def _patched_find_spec(name, package=None):
_blocked = ('tensorflow', 'jax', 'jaxlib')
if name in _blocked or any(name.startswith(b + '.') for b in _blocked):
return None
return _original_find_spec(name, package)
importlib.util.find_spec = _patched_find_spec
import numpy as np
# numpy 2.x 兼容补丁(imgaug 依赖 np.sctypes)
if not hasattr(np, 'sctypes'):
np.sctypes = {
'int': [np.int8, np.int16, np.int32, np.int64],
'uint': [np.uint8, np.uint16, np.uint32, np.uint64],
'float': [np.float16, np.float32, np.float64],
'complex': [np.complex64, np.complex128],
'others': [bool, object, bytes, str, np.void],
}
from module.config import DEVICE, BGE_MODEL_PATH
_model = None
_dinov2_extractor = None
_ocr_engine = None
_bge_tokenizer = None
_bge_model = None
def get_model():
return _model
def get_dinov2():
global _dinov2_extractor
if _dinov2_extractor is None:
load_dinov2()
return _dinov2_extractor
def get_ocr():
global _ocr_engine
if _ocr_engine is None:
load_ocr()
return _ocr_engine
def get_bge():
global _bge_tokenizer, _bge_model
if _bge_tokenizer is None or _bge_model is None:
load_bge()
return _bge_tokenizer, _bge_model
def load_model():
"""MASt3R 已禁用(CPU 部署模式)"""
global _model
print("[INFO] MASt3R 已禁用(CPU 部署,使用 DINOv2-only 模式)")
_model = None
return _model
def load_dinov2():
"""加载 DINOv2 模型(vit_small_patch14_reg4_dinov2)"""
global _dinov2_extractor
if _dinov2_extractor is not None:
return _dinov2_extractor
from module.dinov2_utils import DINOv2Extractor
print("加载 DINOv2 模型...")
try:
_dinov2_extractor = DINOv2Extractor()
if not _dinov2_extractor.is_available:
_dinov2_extractor = None
print("DINOv2 不可用,复杂场景检测将无法工作")
except Exception as e:
_dinov2_extractor = None
print(f"DINOv2 加载异常: {e}")
return _dinov2_extractor
def load_ocr():
"""加载 PaddleOCR 引擎(CPU 模式,原生 Paddle 推理)"""
global _ocr_engine
if _ocr_engine is not None:
return _ocr_engine
print("加载 PaddleOCR...")
try:
from paddleocr import PaddleOCR
_ocr_engine = PaddleOCR(
use_angle_cls=True,
lang='ch',
cpu_threads=2, # 匹配 HF Spaces 免费 2vCPU
ocr_version='PP-OCRv3', # 轻量版,文本场景提速 ~21%,判定结果与 v4 一致
)
print("PaddleOCR (PP-OCRv3) 加载完成")
except Exception as e:
_ocr_engine = None
print(f"PaddleOCR 加载异常: {e}")
print("文本场景检测将不可用")
return _ocr_engine
def load_bge():
"""加载 BGE-small-zh 语义嵌入模型(从 HuggingFace Hub 下载)"""
global _bge_tokenizer, _bge_model
if _bge_tokenizer is not None and _bge_model is not None:
return _bge_tokenizer, _bge_model
print("加载 BGE-small-zh 模型...")
try:
from transformers import AutoTokenizer
from transformers.models.bert.modeling_bert import BertModel
_bge_tokenizer = AutoTokenizer.from_pretrained(BGE_MODEL_PATH)
_bge_model = BertModel.from_pretrained(BGE_MODEL_PATH)
_bge_model.eval()
print("BGE-small-zh 加载完成")
except Exception as e:
_bge_tokenizer = None
_bge_model = None
print(f"BGE-small-zh 加载异常: {e}")
print("文本语义比对将不可用")
return _bge_tokenizer, _bge_model