feat: refactor orientation scoring tasks and enhance batch detection logic

This commit is contained in:
myhloli
2026-06-17 20:00:47 +08:00
parent 87440e9f8a
commit 8fecc190dc
+158 -33
View File
@@ -2,6 +2,7 @@
from PIL import Image
from collections import defaultdict
import inspect
from typing import List, Dict
import cv2
import numpy as np
@@ -129,11 +130,14 @@ class MineruTableOrientationClsModel:
return None
return img[ymin:ymax, xmin:xmax].copy()
def _build_orientation_score_task(self, label: str, img_bgr: np.ndarray) -> Dict:
"""为单个角度构造评分任务,只做 det、抽样和切图,不执行 rec。"""
det_ocr_res = self.ocr_engine.ocr(img_bgr, rec=False)
det_res = det_ocr_res[0] if det_ocr_res else None
sampled_boxes = self._sample_det_boxes(det_res)
def _build_orientation_score_task_from_det_boxes(
self,
label: str,
img_bgr: np.ndarray,
det_boxes,
) -> Dict:
"""根据已有 OCR det 框构造评分任务,复用 0 度门控结果并统一裁图逻辑。"""
sampled_boxes = self._sample_det_boxes(det_boxes)
img_crop_list = []
for box in sampled_boxes:
@@ -149,6 +153,16 @@ class MineruTableOrientationClsModel:
"crop_end": len(img_crop_list),
}
def _build_orientation_score_task(self, label: str, img_bgr: np.ndarray) -> Dict:
"""为单个角度构造评分任务,只做 det、抽样和切图,不执行 rec。"""
det_ocr_res = self.ocr_engine.ocr(img_bgr, rec=False)
det_res = det_ocr_res[0] if det_ocr_res else None
return self._build_orientation_score_task_from_det_boxes(
label,
img_bgr,
det_res,
)
def _build_orientation_score_tasks(self, img_bgr: np.ndarray) -> List[Dict]:
"""为一张表构造 0/90/270 三个角度的评分任务。"""
tasks = []
@@ -182,6 +196,9 @@ class MineruTableOrientationClsModel:
score_by_label = {}
rec_res = rec_res or []
for task in tasks:
if "score" in task:
score_by_label[task["label"]] = task["score"]
continue
crop_start = task.get("crop_start", 0)
crop_end = task.get("crop_end", crop_start + task.get("crop_count", 0))
task_rec_res = rec_res[crop_start:crop_end]
@@ -306,6 +323,24 @@ class MineruTableOrientationClsModel:
)
return resolution_groups
@classmethod
def _collect_orientation_images(cls, imgs: List[Dict]) -> list[Dict]:
"""扁平收集有效表格图,首轮 det 的分桶和 batch 交给 OCR detector 内部处理。"""
orientation_imgs = []
for index, img in enumerate(imgs):
bgr_img = cls._to_bgr_table_image(img)
img_height, img_width = bgr_img.shape[:2]
if img_height <= 0 or img_width <= 0:
continue
orientation_imgs.append(
{
"index": index,
"table_img_bgr": bgr_img,
}
)
return orientation_imgs
@classmethod
def _pad_group_images(
cls,
@@ -327,46 +362,136 @@ class MineruTableOrientationClsModel:
batch_images.append(padded_img)
return batch_images
def _batch_detect_text_boxes(
self,
img_list: list[np.ndarray],
det_batch_size: int,
tqdm_enable: bool = False,
tqdm_desc: str = "OCR-det Predict",
progress_bar=None,
):
"""统一调用 OCR detector batch_predict,并兼容不支持进度参数的测试替身。"""
if not img_list:
return []
max_batch_size = max(1, min(len(img_list), int(det_batch_size)))
batch_predict = self.ocr_engine.text_detector.batch_predict
progress_kwargs = {}
try:
signature = inspect.signature(batch_predict)
params = signature.parameters
except (TypeError, ValueError):
params = {}
if "tqdm_enable" in params:
progress_kwargs["tqdm_enable"] = tqdm_enable
if "tqdm_desc" in params:
progress_kwargs["tqdm_desc"] = tqdm_desc
if "tqdm_progress_bar" in params:
progress_kwargs["tqdm_progress_bar"] = progress_bar
batch_results = batch_predict(img_list, max_batch_size, **progress_kwargs)
if progress_bar is not None and "tqdm_progress_bar" not in progress_kwargs:
progress_bar.update(len(img_list))
return batch_results
def _detect_rotation_candidates(
self,
resolution_groups: Dict[tuple[int, int], list[Dict]],
orientation_imgs: list[Dict],
det_batch_size: int,
resolution_group_stride: int,
tqdm_enable: bool = False,
tqdm_desc: str = "Table orientation",
progress_bar=None,
) -> list[Dict]:
"""对表格批量做 OCR det,并筛选需要进入多角度评分的候选。"""
rotated_imgs = []
for _group_key, group_imgs in resolution_groups.items():
batch_images = self._pad_group_images(group_imgs, resolution_group_stride)
batch_results = self.ocr_engine.text_detector.batch_predict(
batch_images,
max(1, min(len(batch_images), det_batch_size)),
)
batch_images = [img_info["table_img_bgr"] for img_info in orientation_imgs]
batch_results = self._batch_detect_text_boxes(
batch_images,
det_batch_size,
tqdm_enable=tqdm_enable and progress_bar is None,
tqdm_desc=f"{tqdm_desc} det",
progress_bar=progress_bar,
)
for img_info, (dt_boxes, _elapse) in zip(group_imgs, batch_results):
if self._is_rotation_candidate_by_det_boxes(dt_boxes):
rotated_imgs.append(img_info)
if progress_bar is not None:
progress_bar.update(len(group_imgs))
for img_info, (dt_boxes, _elapse) in zip(orientation_imgs, batch_results):
if not self._is_rotation_candidate_by_det_boxes(dt_boxes):
continue
candidate_info = dict(img_info)
candidate_info["gate_det_boxes"] = dt_boxes
rotated_imgs.append(candidate_info)
return rotated_imgs
@staticmethod
def _add_score_task_crops(task: Dict, all_crop_imgs: list[np.ndarray]) -> None:
"""将有效评分任务加入 OCR-rec 输入,不足阈值的任务直接记为 0 分。"""
if task["crop_count"] < ORIENTATION_SCORE_MIN_VALID_RESULTS:
task["score"] = (0.0, 0, 0)
task["crop_start"] = len(all_crop_imgs)
task["crop_end"] = len(all_crop_imgs)
return
crop_start = len(all_crop_imgs)
all_crop_imgs.extend(task["crops"])
crop_end = len(all_crop_imgs)
task["crop_start"] = crop_start
task["crop_end"] = crop_end
def _build_score_tasks_for_candidates(
self,
rotated_imgs: list[Dict],
det_batch_size: int,
progress_bar=None,
) -> tuple[list[tuple[Dict, list[Dict]]], list[np.ndarray]]:
"""为所有旋转候选构造三角度评分任务,并汇总成一次 OCR rec 输入。"""
img_score_tasks = []
all_crop_imgs = []
score_det_images = []
score_det_tasks = []
for img_info in rotated_imgs:
tasks = self._build_orientation_score_tasks(img_info["table_img_bgr"])
for task in tasks:
crop_start = len(all_crop_imgs)
all_crop_imgs.extend(task["crops"])
crop_end = len(all_crop_imgs)
task["crop_start"] = crop_start
task["crop_end"] = crop_end
table_img_bgr = img_info["table_img_bgr"]
tasks = [
self._build_orientation_score_task_from_det_boxes(
"0",
table_img_bgr,
img_info.get("gate_det_boxes"),
)
]
for label in ("90", "270"):
rotated_img = self._rotate_image_by_label(table_img_bgr, label)
task = {
"label": label,
"rotated_img_bgr": rotated_img,
"crops": [],
"crop_count": 0,
"crop_start": 0,
"crop_end": 0,
}
tasks.append(task)
score_det_images.append(rotated_img)
score_det_tasks.append(task)
img_score_tasks.append((img_info, tasks))
score_det_results = self._batch_detect_text_boxes(
score_det_images,
det_batch_size,
)
for task, (dt_boxes, _elapse) in zip(score_det_tasks, score_det_results):
score_task = self._build_orientation_score_task_from_det_boxes(
task["label"],
task["rotated_img_bgr"],
dt_boxes,
)
task.update(score_task)
task.pop("rotated_img_bgr", None)
for _img_info, tasks in img_score_tasks:
for task in tasks:
self._add_score_task_crops(task, all_crop_imgs)
if progress_bar is not None:
progress_bar.update(1)
return img_score_tasks, all_crop_imgs
@@ -395,6 +520,7 @@ class MineruTableOrientationClsModel:
def _score_rotation_candidates(
self,
rotated_imgs: list[Dict],
det_batch_size: int,
tqdm_enable: bool = False,
tqdm_desc: str = "Table orientation",
progress_bar=None,
@@ -406,6 +532,7 @@ class MineruTableOrientationClsModel:
label_by_index = {}
img_score_tasks, all_crop_imgs = self._build_score_tasks_for_candidates(
rotated_imgs,
det_batch_size,
progress_bar=progress_bar,
)
self._extend_progress_total(progress_bar, len(all_crop_imgs))
@@ -437,13 +564,9 @@ class MineruTableOrientationClsModel:
"""
批量预测传入表格图片的旋转角度,只返回角度,不修改输入图片。
"""
RESOLUTION_GROUP_STRIDE = 128
rotate_labels = ["0"] * len(imgs)
resolution_groups = self._collect_orientation_image_groups(
imgs,
RESOLUTION_GROUP_STRIDE,
)
total_images = sum(len(group_imgs) for group_imgs in resolution_groups.values())
orientation_imgs = self._collect_orientation_images(imgs)
total_images = len(orientation_imgs)
progress_bar = None
if tqdm_enable:
progress_bar = tqdm(
@@ -453,9 +576,10 @@ class MineruTableOrientationClsModel:
)
try:
rotated_imgs = self._detect_rotation_candidates(
resolution_groups,
orientation_imgs,
det_batch_size,
RESOLUTION_GROUP_STRIDE,
tqdm_enable=tqdm_enable,
tqdm_desc=tqdm_desc,
progress_bar=progress_bar,
)
self._extend_progress_total(progress_bar, len(rotated_imgs))
@@ -464,6 +588,7 @@ class MineruTableOrientationClsModel:
progress_bar.refresh()
label_by_index = self._score_rotation_candidates(
rotated_imgs,
det_batch_size,
tqdm_enable=tqdm_enable,
tqdm_desc=tqdm_desc,
progress_bar=progress_bar,