Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 48 additions & 7 deletions PPOCRLabel.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,22 @@

__appname__ = "PPOCRLabel"

DEFAULT_RECOGNITION_MODELS = {
"ch": "PP-OCRv5_mobile_rec",
"en": "en_PP-OCRv5_mobile_rec",
"french": "latin_PP-OCRv5_mobile_rec",
"german": "latin_PP-OCRv5_mobile_rec",
"korean": "korean_PP-OCRv5_mobile_rec",
"japan": "PP-OCRv5_server_rec",
}


def getRecognitionModelName(lang, modelName, modelDir):
if modelDir is None and modelName == "PP-OCRv5_mobile_rec":
return DEFAULT_RECOGNITION_MODELS[lang]
return modelName


LABEL_COLORMAP = label_colormap()


Expand Down Expand Up @@ -200,15 +216,18 @@ def get_str(str_id):
self.rec_model_dir = rec_model_dir
self.rec_model_name = rec_model_name
self.cls_model_dir = cls_model_dir
self.model_lang = self.lang if self.lang in DEFAULT_RECOGNITION_MODELS else "ch"
recognition_model_name = getRecognitionModelName(
self.model_lang, self.rec_model_name, self.rec_model_dir
)

params = {
"use_doc_orientation_classify": False,
"use_doc_unwarping": False,
"use_textline_orientation": False,
"device": self.gpu,
"lang": self.lang,
"text_detection_model_name": self.det_model_name,
"text_recognition_model_name": self.rec_model_name,
"text_recognition_model_name": recognition_model_name,
"enable_mkldnn": False,
}

Expand All @@ -221,7 +240,7 @@ def get_str(str_id):

self.ocr = PaddleOCR(**params)
self.text_recognizer = TextRecognition(
model_name=self.rec_model_name,
model_name=recognition_model_name,
model_dir=self.rec_model_dir,
device=self.gpu,
)
Expand Down Expand Up @@ -3574,6 +3593,17 @@ def autolcm(self):
self.comboBox.addItems(
["Chinese & English", "English", "French", "German", "Korean", "Japanese"]
)
model_labels = {
"ch": "Chinese & English",
"en": "English",
"french": "French",
"german": "German",
"korean": "Korean",
"japan": "Japanese",
}
self.comboBox.setCurrentText(
model_labels.get(self.model_lang, "Chinese & English")
)
vbox.addWidget(self.panel)
vbox.addWidget(self.comboBox)
self.dialog = QDialog()
Expand Down Expand Up @@ -3611,15 +3641,18 @@ def modelChoose(self):
if current_text in lg_idx:
choose_lang = lg_idx[current_text]
if hasattr(self, "ocr"):
del self.ocr
rec_model_name = getRecognitionModelName(
choose_lang, self.rec_model_name, self.rec_model_dir
)

params = {
"use_doc_orientation_classify": False,
"use_textline_orientation": False,
"use_doc_unwarping": False,
"text_detection_model_name": self.det_model_name,
"text_recognition_model_name": self.rec_model_name,
"lang": choose_lang,
"text_recognition_model_name": rec_model_name,
"device": self.gpu,
"enable_mkldnn": False,
}
Comment thread
YDLuo-1 marked this conversation as resolved.
if self.det_model_dir is not None:
params["text_detection_model_dir"] = self.det_model_dir
Expand All @@ -3628,7 +3661,14 @@ def modelChoose(self):
if self.cls_model_dir is not None:
params["text_line_orientation_model_dir"] = self.cls_model_dir

self.ocr = PaddleOCR(**params)
selected_ocr = PaddleOCR(**params)
selected_text_recognizer = TextRecognition(
model_name=rec_model_name,
model_dir=self.rec_model_dir,
device=self.gpu,
)
self.ocr = selected_ocr
self.text_recognizer = selected_text_recognizer
if choose_lang in ["ch", "en"]:
if hasattr(self, "table_ocr"):
del self.table_ocr
Expand All @@ -3642,6 +3682,7 @@ def modelChoose(self):
use_region_detection=False,
device=self.gpu,
)
self.model_lang = choose_lang
else:
logger.error("Invalid language selection")
self.dialog.close()
Expand Down
54 changes: 54 additions & 0 deletions tests/test_model_language.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
import ast
import unittest
from pathlib import Path


sourcePath = Path(__file__).parents[1] / "PPOCRLabel.py"
tree = ast.parse(sourcePath.read_text(encoding="utf-8"))
nodes = [
node
for node in tree.body
if (
isinstance(node, ast.Assign)
and any(
isinstance(target, ast.Name) and target.id == "DEFAULT_RECOGNITION_MODELS"
for target in node.targets
)
)
or (isinstance(node, ast.FunctionDef) and node.name == "getRecognitionModelName")
]
namespace = {}
exec(compile(ast.Module(nodes, type_ignores=[]), str(sourcePath), "exec"), namespace)


class ModelLanguageTest(unittest.TestCase):
def test_default_recognition_model_selection(self):
getModelName = namespace["getRecognitionModelName"]
self.assertEqual(
getModelName("en", "PP-OCRv5_mobile_rec", None),
"en_PP-OCRv5_mobile_rec",
)
self.assertEqual(getModelName("en", "custom_rec", None), "custom_rec")
self.assertEqual(
getModelName("en", "PP-OCRv5_mobile_rec", "custom_dir"),
"PP-OCRv5_mobile_rec",
)

def test_language_switch_keeps_mkldnn_disabled(self):
modelChoose = next(
node
for classNode in tree.body
if isinstance(classNode, ast.ClassDef) and classNode.name == "MainWindow"
for node in classNode.body
if isinstance(node, ast.FunctionDef) and node.name == "modelChoose"
)
self.assertTrue(
any(
isinstance(node, ast.Constant) and node.value == "enable_mkldnn"
for node in ast.walk(modelChoose)
)
)


if __name__ == "__main__":
unittest.main()
Loading