diff --git a/mineru/backend/hybrid_flash/hybrid_flash_analyze.py b/mineru/backend/hybrid_flash/hybrid_flash_analyze.py index d56ab108..0bb96822 100644 --- a/mineru/backend/hybrid_flash/hybrid_flash_analyze.py +++ b/mineru/backend/hybrid_flash/hybrid_flash_analyze.py @@ -32,7 +32,10 @@ from mineru.utils.pdfium_guard import ( from mineru.version import __version__ -VLM_VISUAL_LABELS = {"image", "chart"} +VLM_VISUAL_LABELS = {"image", "chart", "seal"} +VLM_VISUAL_TYPE_BY_LABEL = { + "seal": "image", +} VLM_TEXT_LABEL_TO_TYPE = { "abstract": "text", "algorithm": "code", @@ -122,7 +125,7 @@ def _vlm_type_for_layout_det(layout_det: dict, vlm_ocr_enable: bool, table_enabl if label == "display_formula": return "equation" if label in VLM_VISUAL_LABELS: - return label if image_analysis else None + return VLM_VISUAL_TYPE_BY_LABEL.get(label, label) if image_analysis else None if vlm_ocr_enable: return VLM_TEXT_LABEL_TO_TYPE.get(label) return None @@ -201,6 +204,13 @@ def _merge_vlm_sidecar_result(layout_dets: list[dict], sidecar_blocks) -> None: block_content = block.get("content") label = layout_det.get("label") + if label == "seal": + if block_content is not None: + layout_det["text"] = block_content + layout_det["content"] = block_content + layout_det["sub_type"] = "seal" + continue + if label in VLM_VISUAL_LABELS: if block_content is not None: layout_det["content"] = block_content @@ -247,6 +257,19 @@ def _build_analyze_meta(ocr_enable: bool, vlm_ocr_enable: bool) -> dict: } +def _build_pipeline_batch_options(vlm_ocr_enable: bool) -> dict: + """构造 hybrid-flash 调用 pipeline batch 时的识别策略。""" + if vlm_ocr_enable: + return { + "formula_recognition_scope": "none", + "ocr_rec_enable": False, + } + return { + "formula_recognition_scope": "inline_only", + "ocr_rec_enable": True, + } + + def _get_device_for_cleanup(): """获取清理显存用device;测试或轻量环境缺少torch时退回CPU。""" try: @@ -255,14 +278,6 @@ def _get_device_for_cleanup(): return "cpu" -def _ensure_external_layout_api(predictor) -> None: - """确认mineru-vl-utils已提供外部layout抽取接口,避免静默走错VLM layout路径。""" - if not hasattr(predictor, "batch_extract_with_layout"): - raise AttributeError( - "hybrid-flash requires mineru-vl-utils with `MinerUClient.batch_extract_with_layout` support" - ) - - def doc_analyze( pdf_bytes, image_writer: DataWriter | None = None, @@ -281,7 +296,6 @@ def doc_analyze( if predictor is None: predictor = ModelSingleton().get_model(backend, model_path, server_url, **kwargs) predictor = _maybe_enable_serial_execution(predictor, backend) - _ensure_external_layout_api(predictor) device = _get_device_for_cleanup() ocr_enable = _get_ocr_enable(pdf_bytes, parse_method=parse_method) @@ -331,7 +345,8 @@ def doc_analyze( pipeline_inputs, formula_enable=inline_formula_enable, table_enable=False, - formula_recognition_scope="inline_only", + seal_ocr_rec_enable=False, + **_build_pipeline_batch_options(vlm_ocr_enable), ) vlm_blocks_list = [ _build_vlm_layout_blocks( @@ -409,10 +424,6 @@ async def aio_doc_analyze( if predictor is None: predictor = await _get_model_async(backend, model_path, server_url, **kwargs) predictor = _maybe_enable_serial_execution(predictor, backend) - if not hasattr(predictor, "aio_batch_extract_with_layout"): - raise AttributeError( - "hybrid-flash requires mineru-vl-utils with `MinerUClient.aio_batch_extract_with_layout` support" - ) device = _get_device_for_cleanup() ocr_enable = _get_ocr_enable(pdf_bytes, parse_method=parse_method) @@ -443,7 +454,8 @@ async def aio_doc_analyze( pipeline_inputs, formula_enable=inline_formula_enable, table_enable=False, - formula_recognition_scope="inline_only", + seal_ocr_rec_enable=False, + **_build_pipeline_batch_options(vlm_ocr_enable), ) vlm_blocks_list = [ _build_vlm_layout_blocks( diff --git a/mineru/backend/pipeline/batch_analyze.py b/mineru/backend/pipeline/batch_analyze.py index 737dbbda..7cd6ddef 100644 --- a/mineru/backend/pipeline/batch_analyze.py +++ b/mineru/backend/pipeline/batch_analyze.py @@ -58,11 +58,15 @@ class BatchAnalyze: text_ocr_det_batch_enabled: bool | None = None, mask_inline_formula_for_ocr_det: bool = True, formula_recognition_scope: str = "all", + ocr_rec_enable: bool = True, + seal_ocr_rec_enable: bool = True, ): self.batch_ratio = batch_ratio self.formula_enable = get_formula_enable(formula_enable) self.table_enable = get_table_enable(table_enable) self.model_manager = model_manager + self.ocr_rec_enable = ocr_rec_enable + self.seal_ocr_rec_enable = seal_ocr_rec_enable self.enable_ocr_det_batch = enable_ocr_det_batch self.table_ori_cls_batch_enabled = ( enable_ocr_det_batch if table_ori_cls_batch_enabled is None else table_ori_cls_batch_enabled @@ -73,9 +77,9 @@ class BatchAnalyze: self.mask_inline_formula_for_ocr_det = ( get_ocr_det_mask_inline_formula_enable(mask_inline_formula_for_ocr_det) ) - if formula_recognition_scope not in {"all", "inline_only"}: + if formula_recognition_scope not in {"all", "inline_only", "none"}: raise ValueError(f"Unsupported formula_recognition_scope: {formula_recognition_scope}") - # 控制公式识别范围,默认保持pipeline原行为;hybrid-flash只让pipeline处理行内公式。 + # 控制公式识别范围,默认保持pipeline原行为;hybrid-flash可只保留公式det而跳过MFR。 self.formula_recognition_scope = formula_recognition_scope @staticmethod @@ -264,6 +268,9 @@ class BatchAnalyze: return match.expand(replacement) return text + def _formula_recognition_enabled(self) -> bool: + return self.formula_enable and self.formula_recognition_scope != "none" + @classmethod def _extract_table_inline_objects( cls, @@ -362,7 +369,7 @@ class BatchAnalyze: self.model = self.model_manager.get_model( lang=None, - formula_enable=self.formula_enable, + formula_enable=self._formula_recognition_enabled(), table_enable=self.table_enable, ) atom_model_manager = AtomModelSingleton() @@ -381,35 +388,42 @@ class BatchAnalyze: clean_vram(self.model.device, vram_threshold=8) if self.formula_enable: - formula_labels = ["display_formula", "inline_formula"] + all_formula_labels = ["display_formula", "inline_formula"] + formula_labels = all_formula_labels if self.formula_recognition_scope == "inline_only": formula_labels = ["inline_formula"] images_mfd_res = [] for layout_res in images_layout_res: page_formula_res = [] for res in layout_res: - if res.get("label") in formula_labels: + if res.get("label") in all_formula_labels: res.setdefault("latex", "") + if res.get("label") in formula_labels: page_formula_res.append(res) images_mfd_res.append(page_formula_res) - # 公式识别 - images_formula_list = run_mfr_inference( - self.model.mfr_model.batch_predict, - images_mfd_res, - np_images, - batch_size=self.batch_ratio * MFR_BASE_BATCH_SIZE, - ) - mfr_count = 0 - for image_index in range(len(np_images)): - mfr_count += len(images_formula_list[image_index]) - for formula_res, formula_with_latex in zip( - images_mfd_res[image_index], images_formula_list[image_index] - ): - formula_res["latex"] = formula_with_latex.get("latex", "") + if self.formula_recognition_scope != "none": + # 公式识别 + images_formula_list = run_mfr_inference( + self.model.mfr_model.batch_predict, + images_mfd_res, + np_images, + batch_size=self.batch_ratio * MFR_BASE_BATCH_SIZE, + ) + mfr_count = 0 + for image_index in range(len(np_images)): + mfr_count += len(images_formula_list[image_index]) + for formula_res, formula_with_latex in zip( + images_mfd_res[image_index], images_formula_list[image_index] + ): + formula_res["latex"] = formula_with_latex.get("latex", "") - # 清理显存 - clean_vram(self.model.device, vram_threshold=8) + # 清理显存 + clean_vram(self.model.device, vram_threshold=8) + else: + for page_formula_res in images_mfd_res: + for formula_res in page_formula_res: + formula_res["latex"] = "" else: for layout_res in images_layout_res: @@ -418,6 +432,8 @@ class BatchAnalyze: + ocr_should_recognize_text = bool(self.ocr_rec_enable) + ocr_res_list_all_page = [] table_res_list_all_page = [] for index in range(len(np_images)): @@ -768,7 +784,7 @@ class BatchAnalyze: ocr_result_list = get_ocr_result_list( ocr_res, useful_list, - ocr_res_list_dict['ocr_enable'], + ocr_res_list_dict['ocr_enable'] and ocr_should_recognize_text, bgr_image, _lang, ) @@ -812,7 +828,7 @@ class BatchAnalyze: ocr_result_list = get_ocr_result_list( ocr_res, useful_list, - ocr_res_list_dict['ocr_enable'], + ocr_res_list_dict['ocr_enable'] and ocr_should_recognize_text, bgr_image, _lang, ) @@ -901,10 +917,11 @@ class BatchAnalyze: total_processed += len(img_crop_list) seal_ocr_items = [] - for ocr_res_list_dict in ocr_res_list_all_page: - for layout_res_item in ocr_res_list_dict['layout_res']: - if layout_res_item.get("label") == "seal": - seal_ocr_items.append((ocr_res_list_dict, layout_res_item)) + if self.seal_ocr_rec_enable: + for ocr_res_list_dict in ocr_res_list_all_page: + for layout_res_item in ocr_res_list_dict['layout_res']: + if layout_res_item.get("label") == "seal": + seal_ocr_items.append((ocr_res_list_dict, layout_res_item)) seal_ocr_model = None for ocr_res_list_dict, layout_res_item in tqdm(seal_ocr_items, desc="Seal Predict"): @@ -952,7 +969,7 @@ class BatchAnalyze: for ocr_res_list_dict in ocr_res_list_all_page: self._prune_empty_ocr_text_blocks( ocr_res_list_dict["layout_res"], - ocr_res_list_dict["ocr_enable"], + ocr_res_list_dict["ocr_enable"] and self.ocr_rec_enable, ) return images_layout_res diff --git a/mineru/backend/pipeline/pipeline_analyze.py b/mineru/backend/pipeline/pipeline_analyze.py index 3310147a..e8670d18 100644 --- a/mineru/backend/pipeline/pipeline_analyze.py +++ b/mineru/backend/pipeline/pipeline_analyze.py @@ -332,7 +332,9 @@ def batch_image_analyze( images_with_extra_info: List[Tuple[Image.Image, bool, str]], formula_enable=True, table_enable=True, - formula_recognition_scope="all"): + formula_recognition_scope="all", + ocr_rec_enable=True, + seal_ocr_rec_enable=True): from .batch_analyze import BatchAnalyze @@ -384,6 +386,8 @@ def batch_image_analyze( table_enable, enable_ocr_det_batch, formula_recognition_scope=formula_recognition_scope, + ocr_rec_enable=ocr_rec_enable, + seal_ocr_rec_enable=seal_ocr_rec_enable, ) results = batch_model(images_with_extra_info)