mirror of
https://github.com/opendatalab/MinerU.git
synced 2026-09-21 12:42:22 +08:00
feat: enhance table OCR detection with batch processing and new input handling
This commit is contained in:
@@ -59,6 +59,7 @@ class BatchAnalyze:
|
||||
enable_ocr_det_batch: bool = True,
|
||||
table_ori_cls_batch_enabled: bool | None = None,
|
||||
text_ocr_det_batch_enabled: bool | None = None,
|
||||
table_ocr_det_batch_enabled: bool | None = None,
|
||||
mask_inline_formula_for_ocr_det: bool = True,
|
||||
):
|
||||
self.batch_ratio = batch_ratio
|
||||
@@ -72,6 +73,9 @@ class BatchAnalyze:
|
||||
self.text_ocr_det_batch_enabled = (
|
||||
enable_ocr_det_batch if text_ocr_det_batch_enabled is None else text_ocr_det_batch_enabled
|
||||
)
|
||||
self.table_ocr_det_batch_enabled = (
|
||||
enable_ocr_det_batch if table_ocr_det_batch_enabled is None else table_ocr_det_batch_enabled
|
||||
)
|
||||
self.mask_inline_formula_for_ocr_det = (
|
||||
get_ocr_det_mask_inline_formula_enable(mask_inline_formula_for_ocr_det)
|
||||
)
|
||||
@@ -92,6 +96,72 @@ class BatchAnalyze:
|
||||
return bgr_image
|
||||
return self._apply_mask_boxes_to_image(bgr_image, mask_boxes)
|
||||
|
||||
def _build_table_ocr_det_items(self, table_res_list_all_page: list[dict]) -> list[dict]:
|
||||
"""构造表格 OCR-det 输入项,保留原图、遮罩图和后续回填所需信息。"""
|
||||
table_det_items = []
|
||||
for index, table_res_dict in enumerate(table_res_list_all_page):
|
||||
bgr_image = cv2.cvtColor(table_res_dict["table_img"], cv2.COLOR_RGB2BGR)
|
||||
table_inline_objects = (
|
||||
table_res_dict.get("table_inline_objects", [])
|
||||
if self._table_supports_inline_objects(table_res_dict)
|
||||
else []
|
||||
)
|
||||
inline_mask_boxes = [
|
||||
{"bbox": inline_object["table_rel_mask_bbox"]}
|
||||
for inline_object in table_inline_objects
|
||||
]
|
||||
formula_mask_boxes = [
|
||||
{"bbox": inline_object["table_rel_mask_bbox"]}
|
||||
for inline_object in table_inline_objects
|
||||
if inline_object["kind"] == "formula"
|
||||
]
|
||||
det_image = (
|
||||
self._apply_mask_boxes_to_image(bgr_image, inline_mask_boxes)
|
||||
if inline_mask_boxes
|
||||
else bgr_image
|
||||
)
|
||||
table_det_items.append(
|
||||
{
|
||||
"bgr_image": bgr_image,
|
||||
"det_image": det_image,
|
||||
"formula_mask_boxes": formula_mask_boxes,
|
||||
"lang": table_res_dict["lang"],
|
||||
"table_id": index,
|
||||
}
|
||||
)
|
||||
return table_det_items
|
||||
|
||||
def _append_table_ocr_det_result(
|
||||
self,
|
||||
table_det_item: dict,
|
||||
dt_boxes,
|
||||
rec_img_lang_group: dict,
|
||||
) -> None:
|
||||
"""将单表 OCR-det 结果整理成 OCR-rec 输入,并保持表格回填顺序。"""
|
||||
if dt_boxes is None or len(dt_boxes) == 0:
|
||||
return
|
||||
|
||||
ocr_result = dt_boxes
|
||||
formula_mask_boxes = table_det_item["formula_mask_boxes"]
|
||||
if formula_mask_boxes:
|
||||
ocr_result = update_det_boxes(ocr_result, formula_mask_boxes)
|
||||
if not ocr_result:
|
||||
return
|
||||
|
||||
ocr_result = sorted_boxes(ocr_result)
|
||||
for dt_box in ocr_result:
|
||||
dt_box_array = np.asarray(dt_box, dtype=np.float32)
|
||||
rec_img_lang_group.setdefault(table_det_item["lang"], []).append(
|
||||
{
|
||||
"cropped_img": get_rotate_crop_image_for_text_rec(
|
||||
table_det_item["bgr_image"],
|
||||
dt_box_array.copy(),
|
||||
),
|
||||
"dt_box": dt_box_array.copy(),
|
||||
"table_id": table_det_item["table_id"],
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _prune_empty_ocr_text_blocks(layout_res: list[dict], ocr_enable: bool) -> None:
|
||||
if not ocr_enable or not layout_res:
|
||||
@@ -487,7 +557,7 @@ class BatchAnalyze:
|
||||
f"Table classification failed: {e}, using default model"
|
||||
)
|
||||
|
||||
# OCR det 过程,顺序执行
|
||||
# OCR det 过程,默认使用 detector 内部分桶 batch,关闭开关时回退逐表单张路径。
|
||||
rec_img_lang_group = defaultdict(list)
|
||||
det_ocr_engine = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.OCR,
|
||||
@@ -495,46 +565,40 @@ class BatchAnalyze:
|
||||
det_db_unclip_ratio=1.6,
|
||||
enable_merge_det_boxes=False,
|
||||
)
|
||||
for index, table_res_dict in enumerate(
|
||||
tqdm(table_res_list_all_page, desc="Table-ocr det")
|
||||
):
|
||||
bgr_image = cv2.cvtColor(table_res_dict["table_img"], cv2.COLOR_RGB2BGR)
|
||||
table_inline_objects = (
|
||||
table_res_dict.get("table_inline_objects", [])
|
||||
if self._table_supports_inline_objects(table_res_dict)
|
||||
else []
|
||||
)
|
||||
inline_mask_boxes = [
|
||||
{"bbox": inline_object["table_rel_mask_bbox"]}
|
||||
for inline_object in table_inline_objects
|
||||
]
|
||||
formula_mask_boxes = [
|
||||
{"bbox": inline_object["table_rel_mask_bbox"]}
|
||||
for inline_object in table_inline_objects
|
||||
if inline_object["kind"] == "formula"
|
||||
]
|
||||
det_image = (
|
||||
self._apply_mask_boxes_to_image(bgr_image, inline_mask_boxes)
|
||||
if inline_mask_boxes
|
||||
else bgr_image
|
||||
)
|
||||
ocr_result = run_ocr_inference(
|
||||
det_ocr_engine.ocr, det_image, rec=False
|
||||
)[0]
|
||||
if ocr_result and formula_mask_boxes:
|
||||
ocr_result = update_det_boxes(ocr_result, formula_mask_boxes)
|
||||
if ocr_result:
|
||||
ocr_result = sorted_boxes(ocr_result)
|
||||
# 构造需要 OCR 识别的图片字典,包括cropped_img, dt_box, table_id,并按照语言进行分组
|
||||
for dt_box in ocr_result:
|
||||
rec_img_lang_group[table_res_dict["lang"]].append(
|
||||
{
|
||||
"cropped_img": get_rotate_crop_image_for_text_rec(
|
||||
bgr_image, np.asarray(dt_box, dtype=np.float32)
|
||||
),
|
||||
"dt_box": np.asarray(dt_box, dtype=np.float32),
|
||||
"table_id": index,
|
||||
}
|
||||
table_det_items = self._build_table_ocr_det_items(table_res_list_all_page)
|
||||
if self.table_ocr_det_batch_enabled:
|
||||
det_images = [table_det_item["det_image"] for table_det_item in table_det_items]
|
||||
if det_images:
|
||||
det_batch_size = max(
|
||||
1,
|
||||
min(len(det_images), self.batch_ratio * OCR_DET_BASE_BATCH_SIZE),
|
||||
)
|
||||
batch_results = run_ocr_inference(
|
||||
det_ocr_engine.text_detector.batch_predict,
|
||||
det_images,
|
||||
det_batch_size,
|
||||
tqdm_enable=True,
|
||||
tqdm_desc="Table-ocr det",
|
||||
)
|
||||
if len(batch_results) != len(table_det_items):
|
||||
raise ValueError("Table OCR det batch result count mismatch")
|
||||
for table_det_item, (dt_boxes, _) in zip(table_det_items, batch_results):
|
||||
self._append_table_ocr_det_result(
|
||||
table_det_item,
|
||||
dt_boxes,
|
||||
rec_img_lang_group,
|
||||
)
|
||||
else:
|
||||
for table_det_item in tqdm(table_det_items, desc="Table-ocr det"):
|
||||
ocr_result = run_ocr_inference(
|
||||
det_ocr_engine.ocr,
|
||||
table_det_item["det_image"],
|
||||
rec=False,
|
||||
)[0]
|
||||
self._append_table_ocr_det_result(
|
||||
table_det_item,
|
||||
ocr_result,
|
||||
rec_img_lang_group,
|
||||
)
|
||||
|
||||
# OCR rec,按照语言分批处理
|
||||
|
||||
@@ -104,7 +104,6 @@ class ModelPath:
|
||||
pytorch_paddle = "models/OCR/paddleocr_torch"
|
||||
slanet_plus = "models/TabRec/SlanetPlus/slanet-plus.onnx"
|
||||
unet_structure = "models/TabRec/UnetStructure/unet.onnx"
|
||||
paddle_table_cls = "models/TabCls/paddle_table_cls/PP-LCNet_x1_0_table_cls.onnx"
|
||||
|
||||
|
||||
class SplitFlag:
|
||||
|
||||
Reference in New Issue
Block a user