# -*- encoding: utf-8 -*- # @Author: SWHL # @Contact: liekkaskono@163.com from dataclasses import dataclass from pathlib import Path from typing import Any, Dict, List, Sequence, Tuple import gradio as gr import rapidocr from omegaconf import OmegaConf from rapidocr import ( EngineType, LangCls, LangDet, LangRec, ModelType, OCRVersion, RapidOCR, ) OCR_INPUT_FIELDS: Tuple[str, ...] = ( "img_input", "text_score", "box_thresh", "unclip_ratio", "max_side_len", "limit_side_len", "limit_type", "use_dilation", "det_engine", "lang_det", "det_model_type", "det_ocr_version", "cls_engine", "lang_cls", "cls_model_type", "cls_ocr_version", "rec_engine", "lang_rec", "rec_model_type", "rec_ocr_version", "is_word", "use_module", ) DEFAULT_INPUT_VALUES: Dict[str, Any] = { "img_input": None, "text_score": 0.5, "box_thresh": 0.5, "unclip_ratio": 1.6, "max_side_len": 2000, "limit_side_len": 736, "limit_type": "min", "use_dilation": "True", "det_engine": EngineType.ONNXRUNTIME.value, "lang_det": LangDet.CH.value, "det_model_type": ModelType.SMALL.value, "det_ocr_version": OCRVersion.PPOCRV6.value, "cls_engine": EngineType.ONNXRUNTIME.value, "lang_cls": LangCls.CH.value, "cls_model_type": ModelType.MOBILE.value, "cls_ocr_version": OCRVersion.PPOCRV4.value, "rec_engine": EngineType.ONNXRUNTIME.value, "lang_rec": LangRec.CH.value, "rec_model_type": ModelType.SMALL.value, "rec_ocr_version": OCRVersion.PPOCRV6.value, "is_word": "No", "use_module": ["use_det", "use_cls", "use_rec"], } @dataclass(frozen=True) class OCRAppConfig: img_input: Any text_score: float box_thresh: float unclip_ratio: float max_side_len: int limit_side_len: int limit_type: str use_dilation: str det_engine: str lang_det: str det_model_type: str det_ocr_version: str cls_engine: str lang_cls: str cls_model_type: str cls_ocr_version: str rec_engine: str lang_rec: str rec_model_type: str rec_ocr_version: str is_word: str use_module: Sequence[str] @classmethod def from_values(cls, values: Sequence[Any]) -> "OCRAppConfig": if len(values) != len(OCR_INPUT_FIELDS): raise ValueError( f"参数数量不匹配:需要 {len(OCR_INPUT_FIELDS)} 个,实际收到 {len(values)} 个" ) raw_config = dict(zip(OCR_INPUT_FIELDS, values)) for key in ("text_score", "box_thresh", "unclip_ratio"): raw_config[key] = float(raw_config[key]) for key in ("max_side_len", "limit_side_len"): raw_config[key] = int(raw_config[key]) return cls(**raw_config) @property def selected_modules(self) -> Sequence[str]: return self.use_module or [] @property def return_word_box(self) -> bool: return self.is_word == "Yes" @property def use_det(self) -> bool: return "use_det" in self.selected_modules @property def use_cls(self) -> bool: return "use_cls" in self.selected_modules @property def use_rec(self) -> bool: return "use_rec" in self.selected_modules @property def use_dilation_bool(self) -> bool: return self.use_dilation == "True" def to_rapidocr_params(self) -> Dict[str, Any]: return { "Global.max_side_len": self.max_side_len, "Det.engine_type": EngineType(self.det_engine), "Det.lang_type": LangDet(self.lang_det), "Det.model_type": ModelType(self.det_model_type), "Det.ocr_version": OCRVersion(self.det_ocr_version), "Det.use_dilation": self.use_dilation_bool, "Det.limit_side_len": self.limit_side_len, "Det.limit_type": self.limit_type, "Cls.engine_type": EngineType(self.cls_engine), "Cls.lang_type": LangCls(self.lang_cls), "Cls.model_type": ModelType(self.cls_model_type), "Cls.ocr_version": OCRVersion(self.cls_ocr_version), "Rec.engine_type": EngineType(self.rec_engine), "Rec.lang_type": LangRec(self.lang_rec), "Rec.model_type": ModelType(self.rec_model_type), "Rec.ocr_version": OCRVersion(self.rec_ocr_version), } def to_yaml_params(self) -> Dict[str, Any]: return { "Global": { "max_side_len": self.max_side_len, "use_det": self.use_det, "use_cls": self.use_cls, "use_rec": self.use_rec, "return_word_box": self.return_word_box, "text_score": self.text_score, "box_thresh": self.box_thresh, }, "Det": { "engine_type": self.det_engine, "lang_type": self.lang_det, "model_type": self.det_model_type, "ocr_version": self.det_ocr_version, "box_thresh": self.box_thresh, "unclip_ratio": self.unclip_ratio, "use_dilation": self.use_dilation_bool, "limit_side_len": self.limit_side_len, "limit_type": self.limit_type, }, "Cls": { "engine_type": self.cls_engine, "lang_type": self.lang_cls, "model_type": self.cls_model_type, "ocr_version": self.cls_ocr_version, }, "Rec": { "engine_type": self.rec_engine, "lang_type": self.lang_rec, "model_type": self.rec_model_type, "ocr_version": self.rec_ocr_version, }, } def _build_config(values: Sequence[Any]) -> OCRAppConfig: return OCRAppConfig.from_values(values) def _build_example(**overrides: Any) -> List[Any]: example = {**DEFAULT_INPUT_VALUES, **overrides} return [example[field] for field in OCR_INPUT_FIELDS] def _format_ocr_result(ocr_result, config: OCRAppConfig): vis_img = ocr_result.vis() if config.return_word_box: full_word_results = [ word_result for line_word_results in (ocr_result.word_results or []) for word_result in line_word_results ] ocr_txts = [ [idx, txt, score] for idx, (txt, score, _) in enumerate(full_word_results) ] return vis_img, ocr_txts, ocr_result.elapse if not config.use_rec: return vis_img, [], ocr_result.elapse txts = ocr_result.txts or [] scores = ocr_result.scores or [] ocr_txts = [[idx, txt, score] for idx, (txt, score) in enumerate(zip(txts, scores))] return vis_img, ocr_txts, ocr_result.elapse def get_ocr_result(*values): try: config = _build_config(values) ocr_engine = RapidOCR(params=config.to_rapidocr_params()) ocr_result = ocr_engine( config.img_input, use_det=config.use_det, use_cls=config.use_cls, use_rec=config.use_rec, text_score=config.text_score, box_thresh=config.box_thresh, unclip_ratio=config.unclip_ratio, return_word_box=config.return_word_box, ) except Exception as e: err_msg = f"模型加载/识别失败:{str(e)},详细参见 Logs" gr.Warning(err_msg) print(err_msg) return None, [], 0.0 return _format_ocr_result(ocr_result, config) def create_examples() -> List[List[Any]]: return [ _build_example(img_input="images/multi.jpg"), _build_example(img_input="images/ch_en_num.jpg"), _build_example(img_input="images/hand_writen.jpeg"), _build_example(img_input="images/japan.jpg", lang_rec=LangRec.JAPAN.value), _build_example( img_input="images/korean.jpg", det_model_type=ModelType.MOBILE.value, det_ocr_version=OCRVersion.PPOCRV5.value, lang_rec=LangRec.KOREAN.value, rec_model_type=ModelType.MOBILE.value, rec_ocr_version=OCRVersion.PPOCRV5.value, ), ] def export_yaml(*values): config = _build_config(values) default_yaml_path = Path(rapidocr.__file__).parent / "config.yaml" cfg = OmegaConf.load(default_yaml_path) cfg = OmegaConf.merge(cfg, config.to_yaml_params()) save_path = Path(__file__).resolve().parent / "config.yaml" OmegaConf.save(cfg, save_path) return save_path custom_css = """ body {font-family: 'Helvetica Neue', Helvetica;} .gr-button {background-color: #4CAF50; color: white; border: none; padding: 10px 20px; border-radius: 5px;} .gr-button:hover {background-color: #45a049;} .gr-textbox {margin-bottom: 15px;} .example-button {background-color: #1E90FF; color: white; border: none; padding: 8px 15px; border-radius: 5px; margin: 5px;} .example-button:hover {background-color: #FF4500;} .tall-radio .gr-radio-item {padding: 15px 0; min-height: 50px; display: flex; align-items: center;} .tall-radio label {font-size: 16px;} .output-image, .input-image, .image-preview {height: 300px !important} """ with gr.Blocks(title="Rapid⚡OCR Demo", css=custom_css, theme=gr.themes.Soft()) as demo: gr.HTML( """