feat: enhance table OCR detection with batch processing and new input handling

This commit is contained in:
myhloli
2026-06-17 18:16:23 +08:00
parent 57639ebd45
commit 62f09af4e8
2 changed files with 105 additions and 42 deletions
+105 -41
View File
@@ -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,按照语言分批处理
-1
View File
@@ -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: