项目文件夹

文件
wehub-resource-sync 39a086b41b
Build Windows CPU / build (push) Has been cancelled
Build Windows CUDA 10.2 / build (push) Has been cancelled
Build Windows CUDA 11.8 / build (push) Has been cancelled
Build Windows CUDA 12.6 / build (push) Has been cancelled
Build Windows DirectML / build (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:08:08 +08:00

153 行
6.6 KiB
Python

import os
from backend.config import BASE_DIR, config
class PaddleModelConfig:
def __init__(self, hardware_accelerator):
self.hardware_accelerator = hardware_accelerator
# 设置识别语言
self.REC_CHAR_TYPE = config.language.value
# 模型文件目录
self.MODEL_BASE = os.path.join(BASE_DIR, 'models')
# 模型版本 V5
self.MODEL_VERSION = 'V5'
# V5模型默认图形识别的shape为3, 48, 320
self.REC_IMAGE_SHAPE = '3,48,320'
# 初始化模型路径
self.REC_MODEL_PATH = None
self.DET_MODEL_PATH = None
self.DET_MODEL_NAME = None
self.REC_MODEL_NAME = None
# 语言组定义
self.LATIN_LANG = [
'af', 'az', 'bs', 'cs', 'cy', 'da', 'de', 'es', 'et', 'fr', 'ga', 'hr',
'hu', 'id', 'is', 'it', 'ku', 'la', 'lt', 'lv', 'mi', 'ms', 'mt', 'nl',
'no', 'oc', 'pi', 'pl', 'pt', 'ro', 'rs_latin', 'sk', 'sl', 'sq', 'sv',
'sw', 'tl', 'tr', 'uz', 'vi', 'latin', 'german', 'french',
'fi', 'eu', 'gl', 'lb', 'rm', 'ca', 'qu',
]
self.ARABIC_LANG = ['ar', 'fa', 'ug', 'ur', 'ps', 'sd', 'bal']
self.CYRILLIC_LANG = [
'ru', 'rs_cyrillic', 'be', 'bg', 'uk', 'mn', 'abq', 'ady', 'kbd', 'ava',
'dar', 'inh', 'che', 'lbe', 'lez', 'tab', 'cyrillic',
'sr', 'kk', 'ky', 'tg', 'mk', 'tt', 'cv', 'ba', 'mhr', 'mo',
'udm', 'kv', 'os', 'bua', 'xal', 'tyv', 'sah', 'kaa',
]
self.DEVANAGARI_LANG = [
'hi', 'mr', 'ne', 'bh', 'mai', 'ang', 'bho', 'mah', 'sck', 'new', 'gom',
'sa', 'bgc', 'devanagari',
]
self.OTHER_LANG = [
'ch', 'japan', 'korean', 'en', 'ta', 'kn', 'te', 'ka',
'chinese_cht',
]
self.MULTI_LANG = (self.LATIN_LANG + self.ARABIC_LANG + self.CYRILLIC_LANG
+ self.DEVANAGARI_LANG + self.OTHER_LANG)
# 如果设置了识别文本语言类型,则设置为对应的语言
if self.REC_CHAR_TYPE in self.MULTI_LANG:
resolved = self._resolve_models()
if resolved:
self.MODEL_VERSION = 'V5'
self.DET_MODEL_PATH, self.REC_MODEL_PATH, self.DET_MODEL_NAME, self.REC_MODEL_NAME = resolved
def _get_v5_rec_model_name(self, lang):
"""
根据语言获取V5识别模型目录名
参考: https://www.paddleocr.ai/main/version3.x/algorithm/PP-OCRv5/PP-OCRv5_multi_languages.html
"""
if lang in ('ch', 'chinese_cht', 'japan'):
return 'PP-OCRv5_server_rec_infer'
elif lang == 'en':
return 'PP-OCRv5_server_rec_infer'
elif lang == 'korean':
return 'korean_PP-OCRv5_mobile_rec_infer'
elif lang in self.LATIN_LANG:
return 'latin_PP-OCRv5_mobile_rec_infer'
elif lang in self.ARABIC_LANG:
return 'arabic_PP-OCRv5_mobile_rec_infer'
elif lang in self.CYRILLIC_LANG:
return 'cyrillic_PP-OCRv5_mobile_rec_infer'
elif lang in self.DEVANAGARI_LANG:
return 'devanagari_PP-OCRv5_mobile_rec_infer'
elif lang == 'th':
return 'th_PP-OCRv5_mobile_rec_infer'
elif lang == 'el':
return 'el_PP-OCRv5_mobile_rec_infer'
elif lang == 'ta':
return 'ta_PP-OCRv5_mobile_rec_infer'
elif lang == 'te':
return 'te_PP-OCRv5_mobile_rec_infer'
return None
@staticmethod
def _read_model_name_from_yaml(model_dir):
"""从 inference.yml 中读取 Global.model_name"""
yaml_path = os.path.join(model_dir, 'inference.yml')
if not os.path.exists(yaml_path):
return None
try:
with open(yaml_path, 'r', encoding='utf-8') as f:
in_global = False
for line in f:
stripped = line.strip()
if stripped == 'Global:':
in_global = True
continue
if in_global:
if stripped and not stripped.startswith('#') and ':' in stripped:
if stripped.startswith('model_name:'):
return stripped.split(':', 1)[1].strip().strip('"').strip("'")
# 遇到下一个顶级 section 则退出
if stripped and not stripped.startswith('model_name') and not stripped.startswith(' ') and stripped.endswith(':'):
break
except Exception:
pass
return None
def _resolve_models(self):
"""
解析 V5 模型路径,返回 (det_model_path, rec_model_path, det_model_name, rec_model_name) 或 None
"""
v5_base = os.path.join(self.MODEL_BASE, 'V5')
# 快速模式优先使用 mobile 模型,否则使用 server 模型
if config.mode.value == 'fast':
det_model_path = os.path.join(v5_base, 'PP-OCRv5_mobile_det_infer')
if not os.path.exists(det_model_path):
det_model_path = os.path.join(v5_base, 'PP-OCRv5_server_det_infer')
else:
det_model_path = os.path.join(v5_base, 'PP-OCRv5_server_det_infer')
if not os.path.exists(det_model_path):
return None
det_model_name = self._read_model_name_from_yaml(det_model_path)
# 快速模式:中文(简/繁)、英文、日文使用通用 mobile 模型,其他语言使用对应的专用模型
if config.mode.value == 'fast' and self.REC_CHAR_TYPE in ('ch', 'chinese_cht', 'en', 'japan'):
rec_model_path = os.path.join(v5_base, 'PP-OCRv5_mobile_rec_infer')
if os.path.exists(rec_model_path):
rec_model_name = self._read_model_name_from_yaml(rec_model_path)
return det_model_path, rec_model_path, det_model_name, rec_model_name
# mobile 不存在则 fallback 到按语言选择
# 获取识别模型
rec_model_dir_name = self._get_v5_rec_model_name(self.REC_CHAR_TYPE)
if rec_model_dir_name is None:
return None
rec_model_path = os.path.join(v5_base, f'{rec_model_dir_name}_infer'
if not rec_model_dir_name.endswith('_infer')
else rec_model_dir_name)
if not os.path.exists(rec_model_path):
rec_model_path = os.path.join(v5_base, rec_model_dir_name)
if not os.path.exists(rec_model_path):
return None
rec_model_name = self._read_model_name_from_yaml(rec_model_path)
return det_model_path, rec_model_path, det_model_name, rec_model_name