mirror of
https://github.com/opendatalab/MinerU.git
synced 2026-09-24 23:10:23 +08:00
feat: enhance OCR and formula recognition configuration with new flags
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user