diff --git a/mineru/backend/hybrid/hybrid_analyze.py b/mineru/backend/hybrid/hybrid_analyze.py index 89679ed2..46cf05ed 100644 --- a/mineru/backend/hybrid/hybrid_analyze.py +++ b/mineru/backend/hybrid/hybrid_analyze.py @@ -9,7 +9,7 @@ import numpy as np import pypdfium2 as pdfium from loguru import logger from mineru_vl_utils import MinerUClient -from mineru_vl_utils.structs import BlockType +from mineru_vl_utils.structs import BlockType, ContentBlock from tqdm import tqdm from mineru.backend.hybrid.hybrid_model_output_to_middle_json import ( @@ -59,6 +59,44 @@ LAYOUT_TITLE_SPLIT_OVERLAP_THRESHOLD = 0.8 not_extract_list = [item.value for item in NotExtractType] HYBRID_OCR_DET_TEXT_TYPES = set(not_extract_list) +HYBRID_ANALYZE_MODES = {"pro", "flash"} +FLASH_LAYOUT_VISUAL_LABELS = {"image", "chart", "seal"} +FLASH_LAYOUT_LABEL_TO_VLM_TYPE = { + "abstract": BlockType.TEXT, + "algorithm": BlockType.CODE, + "aside_text": BlockType.ASIDE_TEXT, + "content": BlockType.TEXT, + "doc_title": BlockType.TITLE, + "footer": BlockType.FOOTER, + "footer_image": BlockType.FOOTER, + "footnote": BlockType.PAGE_FOOTNOTE, + "formula_number": BlockType.TEXT, + "header": BlockType.HEADER, + "header_image": BlockType.HEADER, + "number": BlockType.PAGE_NUMBER, + "paragraph_title": BlockType.TITLE, + "reference_content": BlockType.REF_TEXT, + "text": BlockType.TEXT, + "vertical_text": BlockType.TEXT, + "figure_title": BlockType.IMAGE_CAPTION, + "vision_footnote": BlockType.IMAGE_FOOTNOTE, + "table": BlockType.TABLE, + "display_formula": BlockType.EQUATION, +} + + +def _validate_hybrid_mode(mode: str) -> str: + """校验 Hybrid 运行模式,避免静默走错解析分支。""" + if mode not in HYBRID_ANALYZE_MODES: + raise ValueError('mode must be "pro" or "flash"') + return mode + + +def _vlm_type_for_flash_layout_label(label: str | None) -> str | None: + """将 pipeline layout 标签映射为 mineru-vl-utils 支持的 VLM 抽取类型。""" + if label in FLASH_LAYOUT_VISUAL_LABELS: + return BlockType.IMAGE + return FLASH_LAYOUT_LABEL_TO_VLM_TYPE.get(label) def _is_hybrid_ocr_det_candidate(block): @@ -281,6 +319,42 @@ def normalize_bbox_to_unit(item, page_width, page_height): return True +def _layout_det_bbox_to_unit(layout_det, page_width, page_height): + """复制并归一化 layout bbox,避免构造 VLM 输入时改动 pipeline 原始结果。""" + bbox = layout_det.get("bbox") + if bbox is None or len(bbox) != 4: + return None + bbox_item = {"bbox": list(bbox)} + if not normalize_bbox_to_unit(bbox_item, page_width, page_height): + return None + return bbox_item["bbox"] + + +def _build_flash_vlm_layout_blocks(layout_dets, page_width, page_height): + """用 pipeline layout 构造 VLM 外部 layout 输入,跳过 VLM 自身 layout 解析。""" + blocks = [] + for layout_det in layout_dets or []: + label = layout_det.get("label") + vlm_type = _vlm_type_for_flash_layout_label(label) + if vlm_type is None: + continue + bbox = _layout_det_bbox_to_unit(layout_det, page_width, page_height) + if bbox is None: + continue + try: + block = ContentBlock( + vlm_type, + bbox, + angle=layout_det.get("angle", 0), + content=layout_det.get("content"), + ) + except AssertionError as exc: + logger.warning(f"Skip invalid Hybrid flash VLM block: {layout_det}, error: {exc}") + continue + blocks.append(block) + return blocks + + def _formula_item_to_pixel_bbox(item): bbox = item.get('bbox') if bbox is not None and len(bbox) == 4: @@ -294,7 +368,7 @@ def _build_inline_formula_inputs(images_layout_res): for layout_res in images_layout_res: page_inline_formula_inputs = [] for res in layout_res: - if res.get('label') not in ['inline_formula', 'display_formula']: + if res.get('label') != 'inline_formula': continue bbox = res.get('bbox') if bbox is None or len(bbox) != 4: @@ -433,13 +507,36 @@ def _predict_layout_for_title_split( ) +def _predict_layout_for_window( + images_pil_list, + language, + inline_formula_enable, + batch_ratio, + vlm_ocr_enable, +): + """为单个处理窗口执行一次 pipeline layout,并返回可复用的小模型实例。""" + hybrid_model_singleton = HybridModelSingleton() + hybrid_pipeline_model = hybrid_model_singleton.get_model( + lang=language, + formula_enable=inline_formula_enable and not vlm_ocr_enable, + ) + images_layout_res = _predict_layout_for_title_split( + hybrid_pipeline_model, + images_pil_list, + batch_ratio, + ) + return images_layout_res, hybrid_pipeline_model + + def _process_ocr_and_formulas( images_pil_list, model_list, - language, inline_formula_enable, _ocr_enable, batch_ratio: int = 1, + *, + images_layout_res, + hybrid_pipeline_model, ): """处理OCR和公式识别""" @@ -450,21 +547,6 @@ def _process_ocr_and_formulas( # 将PIL图片转换为numpy数组 np_images = [np.asarray(pil_image).copy() for pil_image in images_pil_list] - # 获取混合模型实例 - hybrid_model_singleton = HybridModelSingleton() - hybrid_pipeline_model = hybrid_model_singleton.get_model( - lang=language, - formula_enable=inline_formula_enable, - ) - - # 在进行`行内`公式检测和识别前,先将图像中的图片、表格、`行间`公式区域mask掉 - layout_images = mask_image_regions(np_images, model_list) if inline_formula_enable else np_images - images_layout_res = _predict_layout_for_title_split( - hybrid_pipeline_model, - layout_images, - batch_ratio, - ) - if inline_formula_enable: images_mfd_res = _build_inline_formula_inputs(images_layout_res) # 公式识别 @@ -560,38 +642,24 @@ def _process_ocr_and_formulas( if need_ocr_res in page_ocr_res_list: page_ocr_res_list.remove(need_ocr_res) - _apply_layout_title_split( - model_list, - images_layout_res, - [_normalize_page_size(image) for image in images_pil_list], - ) - _normalize_bbox(inline_formula_list, ocr_res_list, images_pil_list) merged_model_list = _merge_page_sidecar_items( model_list, inline_formula_list, ocr_res_list, ) - return merged_model_list, hybrid_pipeline_model + return merged_model_list -def _apply_layout_title_split_for_window( +def _apply_vlm_ocr_det_sidecars_for_window( images_pil_list, model_list, - language, batch_ratio, + *, + images_layout_res, + hybrid_pipeline_model, ): - """为VLM-OCR路径补跑layout小模型,先基于VLM原始title做OCR det,再拆分标题。""" - hybrid_model_singleton = HybridModelSingleton() - hybrid_pipeline_model = hybrid_model_singleton.get_model( - lang=language, - formula_enable=False, - ) - images_layout_res = _predict_layout_for_title_split( - hybrid_pipeline_model, - images_pil_list, - batch_ratio, - ) + """为VLM-OCR路径追加OCR det空文本行和行内公式框sidecar。""" formula_mask_inputs = _build_formula_mask_inputs(images_layout_res) inline_formula_list = _build_inline_formula_det_inputs(images_layout_res) np_images = [np.asarray(pil_image).copy() for pil_image in images_pil_list] @@ -611,12 +679,6 @@ def _apply_layout_title_split_for_window( ocr_res_list, keep_ocr_text=False, ) - _apply_layout_title_split( - model_list, - images_layout_res, - [_normalize_page_size(image) for image in images_pil_list], - ) - return hybrid_pipeline_model def _normalize_bbox( @@ -764,8 +826,10 @@ def doc_analyze( model_path: str | None = None, server_url: str | None = None, image_analysis: bool = True, + mode: str = "pro", **kwargs, ): + mode = _validate_hybrid_mode(mode) client_side_output_generation = bool( kwargs.pop("client_side_output_generation", False) ) @@ -813,22 +877,65 @@ def doc_analyze( ) try: images_pil_list = [image_dict["img_pil"] for image_dict in images_list] + page_sizes = [_normalize_page_size(image) for image in images_pil_list] logger.info( f'Hybrid processing window {window_index + 1}/{total_windows}: ' f'pages {window_start + 1}-{window_end + 1}/{page_count} ' f'({len(images_pil_list)} pages)' ) - if _vlm_ocr_enable: + images_layout_res, hybrid_pipeline_model = _predict_layout_for_window( + images_pil_list, + language, + inline_formula_enable, + batch_ratio, + _vlm_ocr_enable, + ) + if mode == "flash": + vlm_blocks_list = [ + _build_flash_vlm_layout_blocks( + page_layout_res, + pil_img.width, + pil_img.height, + ) + for page_layout_res, pil_img in zip(images_layout_res, images_pil_list) + ] + with predictor_execution_guard(predictor): + window_model_list = predictor.batch_extract_with_layout( + images_pil_list, + vlm_blocks_list, + not_extract_list=None if _vlm_ocr_enable else not_extract_list, + image_analysis=image_analysis, + ) + if _vlm_ocr_enable: + _apply_vlm_ocr_det_sidecars_for_window( + images_pil_list, + window_model_list, + batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, + ) + else: + window_model_list = _process_ocr_and_formulas( + images_pil_list, + window_model_list, + inline_formula_enable, + _ocr_enable, + batch_ratio=batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, + ) + elif _vlm_ocr_enable: with predictor_execution_guard(predictor): window_model_list = predictor.batch_two_step_extract( images=images_pil_list, image_analysis=image_analysis, ) - hybrid_pipeline_model = _apply_layout_title_split_for_window( + _apply_vlm_ocr_det_sidecars_for_window( images_pil_list, window_model_list, - language, batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, ) else: with predictor_execution_guard(predictor): @@ -837,15 +944,21 @@ def doc_analyze( not_extract_list=not_extract_list, image_analysis=image_analysis, ) - window_model_list, hybrid_pipeline_model = _process_ocr_and_formulas( + window_model_list = _process_ocr_and_formulas( images_pil_list, window_model_list, - language, inline_formula_enable, _ocr_enable, batch_ratio=batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, ) + _apply_layout_title_split( + window_model_list, + images_layout_res, + page_sizes, + ) model_list.extend(window_model_list) if progress_bar is None: progress_bar = tqdm(total=page_count, desc="Processing pages") @@ -914,8 +1027,10 @@ async def aio_doc_analyze( model_path: str | None = None, server_url: str | None = None, image_analysis: bool = True, + mode: str = "pro", **kwargs, ): + mode = _validate_hybrid_mode(mode) client_side_output_generation = bool( kwargs.pop("client_side_output_generation", False) ) @@ -962,23 +1077,69 @@ async def aio_doc_analyze( ) try: images_pil_list = [image_dict["img_pil"] for image_dict in images_list] + page_sizes = [_normalize_page_size(image) for image in images_pil_list] logger.info( f'Hybrid processing window {window_index + 1}/{total_windows}: ' f'pages {window_start + 1}-{window_end + 1}/{page_count} ' f'({len(images_pil_list)} pages)' ) - if _vlm_ocr_enable: + images_layout_res, hybrid_pipeline_model = await asyncio.to_thread( + _predict_layout_for_window, + images_pil_list, + language, + inline_formula_enable, + batch_ratio, + _vlm_ocr_enable, + ) + if mode == "flash": + vlm_blocks_list = [ + _build_flash_vlm_layout_blocks( + page_layout_res, + pil_img.width, + pil_img.height, + ) + for page_layout_res, pil_img in zip(images_layout_res, images_pil_list) + ] + async with aio_predictor_execution_guard(predictor): + window_model_list = await predictor.aio_batch_extract_with_layout( + images_pil_list, + vlm_blocks_list, + not_extract_list=None if _vlm_ocr_enable else not_extract_list, + image_analysis=image_analysis, + ) + if _vlm_ocr_enable: + await asyncio.to_thread( + _apply_vlm_ocr_det_sidecars_for_window, + images_pil_list, + window_model_list, + batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, + ) + else: + window_model_list = await asyncio.to_thread( + _process_ocr_and_formulas, + images_pil_list, + window_model_list, + inline_formula_enable, + _ocr_enable, + batch_ratio=batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, + ) + elif _vlm_ocr_enable: async with aio_predictor_execution_guard(predictor): window_model_list = await predictor.aio_batch_two_step_extract( images=images_pil_list, image_analysis=image_analysis, ) - hybrid_pipeline_model = await asyncio.to_thread( - _apply_layout_title_split_for_window, + await asyncio.to_thread( + _apply_vlm_ocr_det_sidecars_for_window, images_pil_list, window_model_list, - language, batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, ) else: async with aio_predictor_execution_guard(predictor): @@ -987,16 +1148,23 @@ async def aio_doc_analyze( not_extract_list=not_extract_list, image_analysis=image_analysis, ) - window_model_list, hybrid_pipeline_model = await asyncio.to_thread( + window_model_list = await asyncio.to_thread( _process_ocr_and_formulas, images_pil_list, window_model_list, - language, inline_formula_enable, _ocr_enable, batch_ratio=batch_ratio, + images_layout_res=images_layout_res, + hybrid_pipeline_model=hybrid_pipeline_model, ) + await asyncio.to_thread( + _apply_layout_title_split, + window_model_list, + images_layout_res, + page_sizes, + ) model_list.extend(window_model_list) if progress_bar is None: progress_bar = tqdm(total=page_count, desc="Processing pages")