diff --git a/mineru/model/table/cls/mineru_table_ori_cls.py b/mineru/model/table/cls/mineru_table_ori_cls.py index f1081197..250224b4 100644 --- a/mineru/model/table/cls/mineru_table_ori_cls.py +++ b/mineru/model/table/cls/mineru_table_ori_cls.py @@ -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,