diff --git a/PPOCRLabel.py b/PPOCRLabel.py index 04f63e5..112f610 100644 --- a/PPOCRLabel.py +++ b/PPOCRLabel.py @@ -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() @@ -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, } @@ -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, ) @@ -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() @@ -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, } if self.det_model_dir is not None: params["text_detection_model_dir"] = self.det_model_dir @@ -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 @@ -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() diff --git a/tests/test_model_language.py b/tests/test_model_language.py new file mode 100644 index 0000000..eb67b92 --- /dev/null +++ b/tests/test_model_language.py @@ -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()