feat: enhance OCR and formula recognition configuration with new flags

This commit is contained in:
myhloli
2026-06-05 00:05:28 +08:00
parent ce44a62c02
commit b97814b69f
3 changed files with 79 additions and 46 deletions
@@ -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(
+45 -28
View File
@@ -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
+5 -1
View File
@@ -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)