From 57639ebd45dd3a48b6c6b998e74be402c2d34d41 Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 15:06:54 +0800 Subject: [PATCH 1/6] feat: optimize batch processing and inference mode for improved performance --- mineru/backend/hybrid/hybrid_analyze.py | 82 ++---- mineru/backend/pipeline/batch_analyze.py | 97 +++---- mineru/model/ocr/pytorch_paddle.py | 2 +- mineru/model/utils/pytorchocr/base_ocr_v20.py | 2 +- .../modeling/heads/rec_multi_head.py | 4 +- .../pytorchocr/postprocess/rec_postprocess.py | 30 +- mineru/model/utils/tools/infer/predict_cls.py | 5 +- mineru/model/utils/tools/infer/predict_det.py | 267 ++++++++++-------- mineru/model/utils/tools/infer/predict_rec.py | 13 +- 9 files changed, 248 insertions(+), 254 deletions(-) diff --git a/mineru/backend/hybrid/hybrid_analyze.py b/mineru/backend/hybrid/hybrid_analyze.py index 96434e7a..8dff3bbf 100644 --- a/mineru/backend/hybrid/hybrid_analyze.py +++ b/mineru/backend/hybrid/hybrid_analyze.py @@ -257,63 +257,39 @@ def ocr_det( bgr_image, det_image, useful_list, adjusted_mfdetrec_res, ocr_res_list[-1] )) - # 按分辨率分组并同时完成padding - RESOLUTION_GROUP_STRIDE = 64 # 32 + batch_images = [crop_info[1] for crop_info in all_cropped_images_info] + det_batch_size = min(len(batch_images), batch_ratio * OCR_DET_BASE_BATCH_SIZE) + batch_results = run_ocr_inference( + hybrid_pipeline_model.ocr_model.text_detector.batch_predict, + batch_images, + det_batch_size, + tqdm_enable=True, + tqdm_desc="OCR-det", + ) - resolution_groups = defaultdict(list) - for crop_info in all_cropped_images_info: - cropped_img = crop_info[1] - h, w = cropped_img.shape[:2] - # 直接计算目标尺寸并用作分组键 - target_h = ((h + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE - target_w = ((w + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE - group_key = (target_h, target_w) - resolution_groups[group_key].append(crop_info) + for crop_info, (dt_boxes, _) in zip(all_cropped_images_info, batch_results): + bgr_image, _det_image, useful_list, adjusted_mfdetrec_res, ocr_page_res_list = crop_info - # 对每个分辨率组进行批处理 - for (target_h, target_w), group_crops in tqdm(resolution_groups.items(), desc="OCR-det"): - # 对所有图像进行padding到统一尺寸 - batch_images = [] - for crop_info in group_crops: - img = crop_info[1] - h, w = img.shape[:2] - # 创建目标尺寸的白色背景 - padded_img = np.ones((target_h, target_w, 3), dtype=np.uint8) * 255 - padded_img[:h, :w] = img - batch_images.append(padded_img) + if dt_boxes is not None and len(dt_boxes) > 0: + # 处理检测框 + dt_boxes_sorted = sorted_boxes(dt_boxes) + dt_boxes_merged = merge_det_boxes(dt_boxes_sorted) if dt_boxes_sorted else [] - # 批处理检测 - det_batch_size = min(len(batch_images), batch_ratio * OCR_DET_BASE_BATCH_SIZE) - batch_results = run_ocr_inference( - hybrid_pipeline_model.ocr_model.text_detector.batch_predict, - batch_images, - det_batch_size, - ) + # 根据公式位置更新检测框 + dt_boxes_final = (update_det_boxes(dt_boxes_merged, adjusted_mfdetrec_res) + if dt_boxes_merged and adjusted_mfdetrec_res + else dt_boxes_merged) - # 处理批处理结果 - for crop_info, (dt_boxes, _) in zip(group_crops, batch_results): - bgr_image, _det_image, useful_list, adjusted_mfdetrec_res, ocr_page_res_list = crop_info - - if dt_boxes is not None and len(dt_boxes) > 0: - # 处理检测框 - dt_boxes_sorted = sorted_boxes(dt_boxes) - dt_boxes_merged = merge_det_boxes(dt_boxes_sorted) if dt_boxes_sorted else [] - - # 根据公式位置更新检测框 - dt_boxes_final = (update_det_boxes(dt_boxes_merged, adjusted_mfdetrec_res) - if dt_boxes_merged and adjusted_mfdetrec_res - else dt_boxes_merged) - - if dt_boxes_final: - ocr_res = [box.tolist() if hasattr(box, 'tolist') else box for box in dt_boxes_final] - ocr_result_list = get_ocr_result_list( - ocr_res, - useful_list, - False, - bgr_image, - hybrid_pipeline_model.lang, - ) - ocr_page_res_list.extend(ocr_result_list) + if dt_boxes_final: + ocr_res = [box.tolist() if hasattr(box, 'tolist') else box for box in dt_boxes_final] + ocr_result_list = get_ocr_result_list( + ocr_res, + useful_list, + False, + bgr_image, + hybrid_pipeline_model.lang, + ) + ocr_page_res_list.extend(ocr_result_list) return ocr_res_list diff --git a/mineru/backend/pipeline/batch_analyze.py b/mineru/backend/pipeline/batch_analyze.py index a0420304..5ea2cba8 100644 --- a/mineru/backend/pipeline/batch_analyze.py +++ b/mineru/backend/pipeline/batch_analyze.py @@ -687,69 +687,48 @@ class BatchAnalyze: lang=lang ) - # 按分辨率分组并同时完成padding - # RESOLUTION_GROUP_STRIDE = 32 - RESOLUTION_GROUP_STRIDE = 64 + batch_images = [crop_info[1] for crop_info in lang_crop_list] + det_batch_size = min( + len(batch_images), self.batch_ratio * OCR_DET_BASE_BATCH_SIZE + ) + batch_results = run_ocr_inference( + ocr_model.text_detector.batch_predict, + batch_images, + det_batch_size, + tqdm_enable=True, + tqdm_desc=f"OCR-det {lang}", + ) - resolution_groups = defaultdict(list) - for crop_info in lang_crop_list: - cropped_img = crop_info[1] - h, w = cropped_img.shape[:2] - # 直接计算目标尺寸并用作分组键 - target_h = ((h + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE - target_w = ((w + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE - group_key = (target_h, target_w) - resolution_groups[group_key].append(crop_info) + for crop_info, (dt_boxes, _) in zip(lang_crop_list, batch_results): + ( + bgr_image, + _det_image, + useful_list, + ocr_res_list_dict, + adjusted_mfdetrec_res, + _lang, + ) = crop_info - # 对每个分辨率组进行批处理 - for (target_h, target_w), group_crops in tqdm(resolution_groups.items(), desc=f"OCR-det {lang}"): - # 对所有图像进行padding到统一尺寸 - batch_images = [] - for crop_info in group_crops: - img = crop_info[1] - h, w = img.shape[:2] - # 创建目标尺寸的白色背景 - padded_img = np.ones((target_h, target_w, 3), dtype=np.uint8) * 255 - padded_img[:h, :w] = img - batch_images.append(padded_img) + if dt_boxes is not None and len(dt_boxes) > 0: + # 处理检测框 + dt_boxes_sorted = sorted_boxes(dt_boxes) + dt_boxes_merged = merge_det_boxes(dt_boxes_sorted) if dt_boxes_sorted else [] - # 批处理检测 - det_batch_size = min(len(batch_images), self.batch_ratio * OCR_DET_BASE_BATCH_SIZE) - batch_results = run_ocr_inference( - ocr_model.text_detector.batch_predict, batch_images, det_batch_size - ) + # 根据公式位置更新检测框 + dt_boxes_final = (update_det_boxes(dt_boxes_merged, adjusted_mfdetrec_res) + if dt_boxes_merged and adjusted_mfdetrec_res + else dt_boxes_merged) - # 处理批处理结果 - for crop_info, (dt_boxes, _) in zip(group_crops, batch_results): - ( - bgr_image, - _det_image, - useful_list, - ocr_res_list_dict, - adjusted_mfdetrec_res, - _lang, - ) = crop_info - - if dt_boxes is not None and len(dt_boxes) > 0: - # 处理检测框 - dt_boxes_sorted = sorted_boxes(dt_boxes) - dt_boxes_merged = merge_det_boxes(dt_boxes_sorted) if dt_boxes_sorted else [] - - # 根据公式位置更新检测框 - dt_boxes_final = (update_det_boxes(dt_boxes_merged, adjusted_mfdetrec_res) - if dt_boxes_merged and adjusted_mfdetrec_res - else dt_boxes_merged) - - if dt_boxes_final: - ocr_res = [box.tolist() if hasattr(box, 'tolist') else box for box in dt_boxes_final] - ocr_result_list = get_ocr_result_list( - ocr_res, - useful_list, - ocr_res_list_dict['ocr_enable'], - bgr_image, - _lang, - ) - ocr_res_list_dict['layout_res'].extend(ocr_result_list) + if dt_boxes_final: + ocr_res = [box.tolist() if hasattr(box, 'tolist') else box for box in dt_boxes_final] + ocr_result_list = get_ocr_result_list( + ocr_res, + useful_list, + ocr_res_list_dict['ocr_enable'], + bgr_image, + _lang, + ) + ocr_res_list_dict['layout_res'].extend(ocr_result_list) # 清理显存 clean_vram(self.model.device, vram_threshold=8) diff --git a/mineru/model/ocr/pytorch_paddle.py b/mineru/model/ocr/pytorch_paddle.py index 2cf86dc1..d8f09eba 100644 --- a/mineru/model/ocr/pytorch_paddle.py +++ b/mineru/model/ocr/pytorch_paddle.py @@ -245,7 +245,7 @@ class PytorchPaddleOCR(TextSystem): logger.debug("no valid image provided") return None, None - ori_im = img.copy() + ori_im = img dt_boxes, elapse = self.text_detector(img) if dt_boxes is None: diff --git a/mineru/model/utils/pytorchocr/base_ocr_v20.py b/mineru/model/utils/pytorchocr/base_ocr_v20.py index d4177bd0..96af9ef6 100644 --- a/mineru/model/utils/pytorchocr/base_ocr_v20.py +++ b/mineru/model/utils/pytorchocr/base_ocr_v20.py @@ -107,6 +107,6 @@ class BaseOCRV20: # print('model is loaded: {}'.format(weights_path)) def inference(self, inputs): - with torch.no_grad(): + with torch.inference_mode(): infer = self.net(inputs) return infer diff --git a/mineru/model/utils/pytorchocr/modeling/heads/rec_multi_head.py b/mineru/model/utils/pytorchocr/modeling/heads/rec_multi_head.py index 47cebc10..ebcabf3a 100644 --- a/mineru/model/utils/pytorchocr/modeling/heads/rec_multi_head.py +++ b/mineru/model/utils/pytorchocr/modeling/heads/rec_multi_head.py @@ -1,5 +1,4 @@ # Copyright (c) Opendatalab. All rights reserved. -import torch.nn.functional as F from torch import nn from ..necks.rnn import EncoderWithLightSVTR, Im2Seq, SequenceEncoder @@ -71,7 +70,8 @@ class MultiHead(nn.Module): ctc_encoder = ctc_encoder.squeeze(dim=2).permute(0, 2, 1) predicts = self.head(ctc_encoder) if not self.training: - predicts = F.softmax(predicts, dim=2) + # 推理阶段保留 raw logits,交给 CTC 解码端按需计算 max 概率,避免整块 softmax 矩阵搬运。 + return {"ctc_logits": predicts, "ctc_use_raw_logits": True} return predicts ctc_encoder = self.ctc_encoder(x) return self.ctc_head(ctc_encoder) diff --git a/mineru/model/utils/pytorchocr/postprocess/rec_postprocess.py b/mineru/model/utils/pytorchocr/postprocess/rec_postprocess.py index e4cd4b06..d1400455 100644 --- a/mineru/model/utils/pytorchocr/postprocess/rec_postprocess.py +++ b/mineru/model/utils/pytorchocr/postprocess/rec_postprocess.py @@ -181,11 +181,33 @@ class CTCLabelDecode(BaseRecLabelDecode): super(CTCLabelDecode, self).__init__(character_dict_path, use_space_char) + def _decode_raw_logits(self, preds): + """从 raw logits 直接计算 CTC argmax 和 max softmax 概率,避免完整 softmax。""" + logits = preds["ctc_logits"] + if torch.is_tensor(logits): + preds_idx = logits.argmax(dim=2) + max_logits = logits.amax(dim=2) + preds_prob = torch.exp(max_logits - torch.logsumexp(logits, dim=2)) + return preds_prob.float().cpu().numpy(), preds_idx.cpu().numpy() + + logits = np.asarray(logits) + preds_idx = logits.argmax(axis=2) + max_logits = logits.max(axis=2) + stable_logits = logits - max_logits[:, :, None] + logsumexp = max_logits + np.log(np.exp(stable_logits).sum(axis=2)) + preds_prob = np.exp(max_logits - logsumexp).astype(np.float32) + return preds_prob, preds_idx + def __call__(self, preds, label=None, return_word_box=False, *args, **kwargs): - preds_prob, preds_idx = preds.max(axis=2) + if isinstance(preds, dict) and preds.get("ctc_use_raw_logits"): + preds_prob, preds_idx = self._decode_raw_logits(preds) + else: + preds_prob, preds_idx = preds.max(axis=2) + preds_idx = preds_idx.cpu().numpy() + preds_prob = preds_prob.float().cpu().numpy() text = self.decode( - preds_idx.cpu().numpy(), - preds_prob.float().cpu().numpy(), + preds_idx, + preds_prob, is_remove_duplicate=True, return_word_box=return_word_box, ) @@ -787,4 +809,4 @@ class CANLabelDecode(BaseRecLabelDecode): if label is None: return text label = self.decode(label) - return text, label \ No newline at end of file + return text, label diff --git a/mineru/model/utils/tools/infer/predict_cls.py b/mineru/model/utils/tools/infer/predict_cls.py index 5e08151d..f4ca6444 100644 --- a/mineru/model/utils/tools/infer/predict_cls.py +++ b/mineru/model/utils/tools/infer/predict_cls.py @@ -86,11 +86,10 @@ class TextClassifier(BaseOCRV20): norm_img = self.resize_norm_img(img_list[indices[ino]]) norm_img = norm_img[np.newaxis, :] norm_img_batch.append(norm_img) - norm_img_batch = np.concatenate(norm_img_batch) - norm_img_batch = norm_img_batch.copy() + norm_img_batch = np.ascontiguousarray(np.concatenate(norm_img_batch)) starttime = time.time() - with torch.no_grad(): + with torch.inference_mode(): inp = torch.from_numpy(norm_img_batch) inp = inp.to(self.device) inp = self._to_inference_dtype(inp) diff --git a/mineru/model/utils/tools/infer/predict_det.py b/mineru/model/utils/tools/infer/predict_det.py index f846c740..ec1e84e2 100644 --- a/mineru/model/utils/tools/infer/predict_det.py +++ b/mineru/model/utils/tools/infer/predict_det.py @@ -1,9 +1,11 @@ # Copyright (c) Opendatalab. All rights reserved. import sys +from collections import defaultdict import numpy as np import time import torch +from tqdm import tqdm from ...pytorchocr.base_ocr_v20 import BaseOCRV20 from . import pytorchocr_utility as utility from ...pytorchocr.data import create_operators, transform @@ -124,6 +126,90 @@ class TextDetector(BaseOCRV20): if hasattr(module, 'rep'): module.rep() + def _preprocess_det_image(self, img): + """执行 OCR-det 单图预处理,并保留后处理需要的原始尺寸信息。""" + data = {'image': img} + data = transform(data, self.preprocess_op) + if data is None: + return None + + img_processed, shape_list = data + if img_processed is None: + return None + + return np.ascontiguousarray(img_processed), shape_list, img.shape + + def _build_det_preds(self, outputs): + """将模型输出统一转换为后处理需要的 float32 numpy 结构。""" + preds = {} + if self.det_algorithm == "EAST": + preds['f_geo'] = outputs['f_geo'].float().cpu().numpy() + preds['f_score'] = outputs['f_score'].float().cpu().numpy() + elif self.det_algorithm == 'SAST': + preds['f_border'] = outputs['f_border'].float().cpu().numpy() + preds['f_score'] = outputs['f_score'].float().cpu().numpy() + preds['f_tco'] = outputs['f_tco'].float().cpu().numpy() + preds['f_tvo'] = outputs['f_tvo'].float().cpu().numpy() + elif self.det_algorithm in ['DB', 'PSE', 'DB++']: + preds['maps'] = outputs['maps'].float().cpu().numpy() + elif self.det_algorithm == 'FCE': + for i, (_k, output) in enumerate(outputs.items()): + preds['level_{}'.format(i)] = output.float().cpu().numpy() + else: + raise NotImplementedError + return preds + + def _postprocess_det_batch(self, preds, batch_shapes, ori_shapes): + """对完整 batch 执行一次 OCR-det 后处理,再逐张裁剪过滤检测框。""" + post_results = self.postprocess_op(preds, batch_shapes) + batch_results = [] + for post_result, ori_shape in zip(post_results, ori_shapes): + dt_boxes = post_result['points'] + dt_boxes = self._filter_det_res(dt_boxes, ori_shape) + batch_results.append(dt_boxes) + return batch_results + + def _batch_process_preprocessed(self, batch_items): + """对已经完成预处理且形状一致的图片执行批量推理。""" + starttime = time.time() + if not batch_items: + return [], 0 + + batch_data = [item[1] for item in batch_items] + batch_shapes = [item[2] for item in batch_items] + ori_shapes = [item[3] for item in batch_items] + + try: + batch_tensor = np.ascontiguousarray(np.stack(batch_data, axis=0)) + batch_shapes = np.stack(batch_shapes, axis=0) + except Exception: + batch_results = [] + for _index, img_processed, shape_list, ori_shape in batch_items: + single_tensor = np.expand_dims(np.ascontiguousarray(img_processed), axis=0) + single_shape = np.expand_dims(shape_list, axis=0) + with torch.inference_mode(): + inp = torch.from_numpy(single_tensor) + inp = inp.to(self.device) + inp = self._to_inference_dtype(inp) + outputs = self.net(inp) + preds = self._build_det_preds(outputs) + dt_boxes = self._postprocess_det_batch(preds, single_shape, [ori_shape])[0] + batch_results.append((dt_boxes, 0)) + return batch_results, time.time() - starttime + + with torch.inference_mode(): + inp = torch.from_numpy(batch_tensor) + inp = inp.to(self.device) + inp = self._to_inference_dtype(inp) + outputs = self.net(inp) + + preds = self._build_det_preds(outputs) + dt_boxes_batch = self._postprocess_det_batch(preds, batch_shapes, ori_shapes) + total_elapse = time.time() - starttime + batch_elapse = total_elapse / len(batch_items) + batch_results = [(dt_boxes, batch_elapse) for dt_boxes in dt_boxes_batch] + return batch_results, total_elapse + def _should_only_clip_det_res(self): if self.det_algorithm == "SAST" and getattr(self, "det_sast_polygon", False): return True @@ -149,93 +235,34 @@ class TextDetector(BaseOCRV20): """ starttime = time.time() - # 预处理所有图像 - batch_data = [] - batch_shapes = [] - ori_imgs = [] - - for img in img_list: - ori_im = img.copy() - ori_imgs.append(ori_im) - - data = {'image': img} - data = transform(data, self.preprocess_op) - if data is None: - # 如果预处理失败,返回空结果 + batch_items = [] + for index, img in enumerate(img_list): + preprocessed = self._preprocess_det_image(img) + if preprocessed is None: return [(None, 0) for _ in img_list], 0 + img_processed, shape_list, ori_shape = preprocessed + batch_items.append((index, img_processed, shape_list, ori_shape)) - img_processed, shape_list = data - batch_data.append(img_processed) - batch_shapes.append(shape_list) + batch_results, _elapsed = self._batch_process_preprocessed(batch_items) + return batch_results, time.time() - starttime - # 堆叠成批处理张量 - try: - batch_tensor = np.stack(batch_data, axis=0) - batch_shapes = np.stack(batch_shapes, axis=0) - except Exception as e: - # 如果堆叠失败,回退到逐个处理 - batch_results = [] - for img in img_list: - dt_boxes, elapse = self.__call__(img) - batch_results.append((dt_boxes, elapse)) - return batch_results, time.time() - starttime - - # 批处理推理 - with torch.no_grad(): - inp = torch.from_numpy(batch_tensor) - inp = inp.to(self.device) - inp = self._to_inference_dtype(inp) - outputs = self.net(inp) - - # 处理输出 - preds = {} - if self.det_algorithm == "EAST": - preds['f_geo'] = outputs['f_geo'].float().cpu().numpy() - preds['f_score'] = outputs['f_score'].float().cpu().numpy() - elif self.det_algorithm == 'SAST': - preds['f_border'] = outputs['f_border'].float().cpu().numpy() - preds['f_score'] = outputs['f_score'].float().cpu().numpy() - preds['f_tco'] = outputs['f_tco'].float().cpu().numpy() - preds['f_tvo'] = outputs['f_tvo'].float().cpu().numpy() - elif self.det_algorithm in ['DB', 'PSE', 'DB++']: - preds['maps'] = outputs['maps'].float().cpu().numpy() - elif self.det_algorithm == 'FCE': - for i, (k, output) in enumerate(outputs.items()): - preds['level_{}'.format(i)] = output.float().cpu().numpy() - else: - raise NotImplementedError - - # 后处理每个图像的结果 - batch_results = [] - total_elapse = time.time() - starttime - - for i in range(len(img_list)): - # 提取单个图像的预测结果 - single_preds = {} - for key, value in preds.items(): - if isinstance(value, np.ndarray): - single_preds[key] = value[i:i + 1] # 保持批次维度 - else: - single_preds[key] = value - - # 后处理 - post_result = self.postprocess_op(single_preds, batch_shapes[i:i + 1]) - dt_boxes = post_result[0]['points'] - - # 过滤和裁剪检测框 - dt_boxes = self._filter_det_res(dt_boxes, ori_imgs[i].shape) - - batch_results.append((dt_boxes, total_elapse / len(img_list))) - - return batch_results, total_elapse - - def batch_predict(self, img_list, max_batch_size=8): + def batch_predict( + self, + img_list, + max_batch_size=8, + tqdm_enable=False, + tqdm_desc="OCR-det Predict", + tqdm_progress_bar=None, + ): """ 批处理预测方法,支持多张图像同时检测 Args: img_list: 图像列表 max_batch_size: 最大批处理大小 + tqdm_enable: 是否显示内部 OCR-det 进度条 + tqdm_desc: 内部 OCR-det 进度条描述 + tqdm_progress_bar: 外部复用进度条,传入时不在本方法内关闭 Returns: batch_results: 批处理结果列表,每个元素为(dt_boxes, elapse) @@ -243,14 +270,38 @@ class TextDetector(BaseOCRV20): if not img_list: return [] - batch_results = [] + progress_bar = tqdm_progress_bar + should_close_progress = False + if progress_bar is None: + progress_bar = tqdm(total=len(img_list), desc=tqdm_desc, disable=not tqdm_enable) + should_close_progress = True - # 分批处理 - for i in range(0, len(img_list), max_batch_size): - batch_imgs = img_list[i:i + max_batch_size] - # assert尺寸一致 - batch_dt_boxes, batch_elapse = self._batch_process_same_size(batch_imgs) - batch_results.extend(batch_dt_boxes) + max_batch_size = max(1, int(max_batch_size)) + batch_results = [(None, 0)] * len(img_list) + grouped_items = defaultdict(list) + + try: + for index, img in enumerate(img_list): + preprocessed = self._preprocess_det_image(img) + if preprocessed is None: + progress_bar.update(1) + continue + img_processed, shape_list, ori_shape = preprocessed + grouped_items[img_processed.shape].append( + (index, img_processed, shape_list, ori_shape) + ) + + for group_items in grouped_items.values(): + for i in range(0, len(group_items), max_batch_size): + batch_items = group_items[i:i + max_batch_size] + group_results, _batch_elapse = self._batch_process_preprocessed(batch_items) + for batch_item, batch_result in zip(batch_items, group_results): + original_index = batch_item[0] + batch_results[original_index] = batch_result + progress_bar.update(len(batch_items)) + finally: + if should_close_progress: + progress_bar.close() return batch_results @@ -310,43 +361,11 @@ class TextDetector(BaseOCRV20): return dt_boxes_new def __call__(self, img): - ori_shape = img.shape - data = {'image': img} - data = transform(data, self.preprocess_op) - img, shape_list = data - if img is None: + preprocessed = self._preprocess_det_image(img) + if preprocessed is None: return None, 0 - img = np.expand_dims(img, axis=0) - shape_list = np.expand_dims(shape_list, axis=0) - img = img.copy() - starttime = time.time() - - with torch.no_grad(): - inp = torch.from_numpy(img) - inp = inp.to(self.device) - inp = self._to_inference_dtype(inp) - outputs = self.net(inp) - - preds = {} - if self.det_algorithm == "EAST": - preds['f_geo'] = outputs['f_geo'].float().cpu().numpy() - preds['f_score'] = outputs['f_score'].float().cpu().numpy() - elif self.det_algorithm == 'SAST': - preds['f_border'] = outputs['f_border'].float().cpu().numpy() - preds['f_score'] = outputs['f_score'].float().cpu().numpy() - preds['f_tco'] = outputs['f_tco'].float().cpu().numpy() - preds['f_tvo'] = outputs['f_tvo'].float().cpu().numpy() - elif self.det_algorithm in ['DB', 'PSE', 'DB++']: - preds['maps'] = outputs['maps'].float().cpu().numpy() - elif self.det_algorithm == 'FCE': - for i, (k, output) in enumerate(outputs.items()): - preds['level_{}'.format(i)] = output.float().cpu().numpy() - else: - raise NotImplementedError - - post_result = self.postprocess_op(preds, shape_list) - dt_boxes = post_result[0]['points'] - dt_boxes = self._filter_det_res(dt_boxes, ori_shape) - - elapse = time.time() - starttime - return dt_boxes, elapse + img_processed, shape_list, ori_shape = preprocessed + batch_results, _elapsed = self._batch_process_preprocessed( + [(0, img_processed, shape_list, ori_shape)] + ) + return batch_results[0] diff --git a/mineru/model/utils/tools/infer/predict_rec.py b/mineru/model/utils/tools/infer/predict_rec.py index 944671bc..45fbcd77 100644 --- a/mineru/model/utils/tools/infer/predict_rec.py +++ b/mineru/model/utils/tools/infer/predict_rec.py @@ -364,8 +364,7 @@ class TextRecognizer(BaseOCRV20): max_wh_ratio) norm_img = norm_img[np.newaxis, :] norm_img_batch.append(norm_img) - norm_img_batch = np.concatenate(norm_img_batch) - norm_img_batch = norm_img_batch.copy() + norm_img_batch = np.ascontiguousarray(np.concatenate(norm_img_batch)) if self.rec_algorithm == "SRN": starttime = time.time() @@ -376,7 +375,7 @@ class TextRecognizer(BaseOCRV20): gsrm_slf_attn_bias2_list = np.concatenate( gsrm_slf_attn_bias2_list) - with torch.no_grad(): + with torch.inference_mode(): inp = torch.from_numpy(norm_img_batch) encoder_word_pos_inp = torch.from_numpy(encoder_word_pos_list) gsrm_word_pos_inp = torch.from_numpy(gsrm_word_pos_list) @@ -407,7 +406,7 @@ class TextRecognizer(BaseOCRV20): # valid_ratios, # ] - with torch.no_grad(): + with torch.inference_mode(): inp = torch.from_numpy(norm_img_batch) inp = inp.to(self.device) inp = self._to_inference_dtype(inp) @@ -422,7 +421,7 @@ class TextRecognizer(BaseOCRV20): inp = [torch.from_numpy(e_i) for e_i in inputs] inp = [e_i.to(self.device) for e_i in inp] inp = [self._to_inference_dtype(e_i) for e_i in inp] - with torch.no_grad(): + with torch.inference_mode(): outputs = self.net(inp) outputs = [v.cpu().numpy() for k, v in enumerate(outputs)] @@ -431,13 +430,13 @@ class TextRecognizer(BaseOCRV20): else: starttime = time.time() - with torch.no_grad(): + with torch.inference_mode(): inp = torch.from_numpy(norm_img_batch) inp = inp.to(self.device) inp = self._to_inference_dtype(inp) preds = self.net(inp) - with torch.no_grad(): + with torch.inference_mode(): rec_result = self.postprocess_op(preds) for rno in range(len(rec_result)): From 62f09af4e8801658524a147cae3cdae0d87de217 Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 18:16:23 +0800 Subject: [PATCH 2/6] feat: enhance table OCR detection with batch processing and new input handling --- mineru/backend/pipeline/batch_analyze.py | 146 ++++++++++++++++------- mineru/utils/enum_class.py | 1 - 2 files changed, 105 insertions(+), 42 deletions(-) diff --git a/mineru/backend/pipeline/batch_analyze.py b/mineru/backend/pipeline/batch_analyze.py index 5ea2cba8..82d7c15f 100644 --- a/mineru/backend/pipeline/batch_analyze.py +++ b/mineru/backend/pipeline/batch_analyze.py @@ -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,按照语言分批处理 diff --git a/mineru/utils/enum_class.py b/mineru/utils/enum_class.py index d666903c..b50faded 100644 --- a/mineru/utils/enum_class.py +++ b/mineru/utils/enum_class.py @@ -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: From 209ada00a103781871c11ca4b999e13293bb69fa Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 18:37:42 +0800 Subject: [PATCH 3/6] feat: add new model path for paddle table classification --- mineru/utils/enum_class.py | 1 + 1 file changed, 1 insertion(+) diff --git a/mineru/utils/enum_class.py b/mineru/utils/enum_class.py index b50faded..d666903c 100644 --- a/mineru/utils/enum_class.py +++ b/mineru/utils/enum_class.py @@ -104,6 +104,7 @@ 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: From 2b8ee2df65c410dabc3b974afe2fcaa906fde814 Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 19:32:43 +0800 Subject: [PATCH 4/6] feat: add blank cell result filtering for OCR to improve accuracy --- mineru/model/table/rec/unet_table/main.py | 50 +++++++++++++---------- 1 file changed, 28 insertions(+), 22 deletions(-) diff --git a/mineru/model/table/rec/unet_table/main.py b/mineru/model/table/rec/unet_table/main.py index 3bb779b4..497bb6f6 100644 --- a/mineru/model/table/rec/unet_table/main.py +++ b/mineru/model/table/rec/unet_table/main.py @@ -28,6 +28,21 @@ from .utils_table_recover import ( gather_ocr_list_by_row, ) +BLANK_CELL_REC_DROP_TEXTS = { + "1", + "一", + "—", + "口", + "■", + "(204号", + "(20", + "(2", + "(2号", + "(20号", + "号", + "(204", +} + @dataclass class WiredTableInput: @@ -155,20 +170,18 @@ class WiredTableRecognition: ) return res - # def fill_blank_rec( - # self, - # img: np.ndarray, - # sorted_polygons: np.ndarray, - # cell_box_map: Dict[int, List[str]], - # ) -> Dict[int, List[Any]]: - # """找到poly对应为空的框,尝试将直接将poly框直接送到识别中""" - # for i in range(sorted_polygons.shape[0]): - # if cell_box_map.get(i): - # continue - # box = sorted_polygons[i] - # cell_box_map[i] = [[box, "", 1]] - # continue - # return cell_box_map + @staticmethod + def _should_drop_blank_cell_rec_result(text: str, score) -> bool: + """判断空单元格二次 OCR-rec 结果是否应作为噪声过滤。""" + try: + if float(score) < 0.6: + return True + except (TypeError, ValueError): + return True + + normalized_text = "" if text is None else str(text).strip() + return not normalized_text or normalized_text in BLANK_CELL_REC_DROP_TEXTS + def fill_blank_rec( self, img: np.ndarray, @@ -209,13 +222,6 @@ class WiredTableRecognition: if len(img_crop_list) > 0: # 进行ocr识别 ocr_result = self.ocr_engine.ocr(img_crop_list, det=False) - # ocr_result = [[]] - # for crop_img in img_crop_list: - # tmp_ocr_result = self.ocr_engine.ocr(crop_img) - # if tmp_ocr_result[0] and len(tmp_ocr_result[0]) > 0 and isinstance(tmp_ocr_result[0], list) and len(tmp_ocr_result[0][0]) == 2: - # ocr_result[0].append(tmp_ocr_result[0][0][1]) - # else: - # ocr_result[0].append(("", 0.0)) if not ocr_result or not isinstance(ocr_result, list) or len(ocr_result) == 0: logger.warning("OCR engine returned no results or invalid result for image crops.") @@ -231,7 +237,7 @@ class WiredTableRecognition: # 处理ocr结果 ocr_text, ocr_score = ocr_res # logger.debug(f"OCR result for box {i}: {ocr_text} with score {ocr_score}") - if ocr_score < 0.6 or ocr_text in ['1','口','■','(204号', '(20', '(2', '(2号', '(20号', '号', '(204']: + if self._should_drop_blank_cell_rec_result(ocr_text, ocr_score): # logger.warning(f"Low confidence OCR result for box {i}: {ocr_text} with score {ocr_score}") box = sorted_polygons[i] cell_box_map[i] = [[box, "", 0.1]] From 87440e9f8a693006ed92da6b27e5d065af77f9d4 Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 19:37:27 +0800 Subject: [PATCH 5/6] feat: add debug logging for OCR orientation recognition scores --- .../model/table/cls/mineru_table_ori_cls.py | 23 ++++++++++++++++++- 1 file changed, 22 insertions(+), 1 deletion(-) diff --git a/mineru/model/table/cls/mineru_table_ori_cls.py b/mineru/model/table/cls/mineru_table_ori_cls.py index e31da4fb..f1081197 100644 --- a/mineru/model/table/cls/mineru_table_ori_cls.py +++ b/mineru/model/table/cls/mineru_table_ori_cls.py @@ -5,6 +5,7 @@ from collections import defaultdict from typing import List, Dict import cv2 import numpy as np +from loguru import logger from tqdm import tqdm @@ -16,7 +17,7 @@ ROTATED_TEXT_MIN_BOXES = 3 ORIENTATION_SCORE_MAX_SAMPLE_BOXES = 18 ORIENTATION_SCORE_MIN_VALID_RESULTS = 5 ORIENTATION_ZERO_SCORE_PRIORITY_THRESHOLD = 0.9 -ORIENTATION_SCORE_TIE_THRESHOLD = 0.1 +ORIENTATION_SCORE_TIE_THRESHOLD = 0.08 ORIENTATION_SCORE_LABELS = ("0", "90", "270") @@ -187,6 +188,24 @@ class MineruTableOrientationClsModel: score_by_label[task["label"]] = self._score_rec_results(task_rec_res) return score_by_label + @staticmethod + def _debug_log_orientation_rec_scores( + table_index: int, + score_by_label: Dict[str, tuple[float, int, int]], + ) -> None: + """输出单张表格各旋转候选的 OCR-rec 分数,便于排查误旋转。""" + score_parts = [] + for label in ORIENTATION_SCORE_LABELS: + score, valid_count, char_count = score_by_label.get(label, (0.0, 0, 0)) + score_parts.append( + f"{label} score={float(score):.4f} " + f"valid_count={valid_count} char_count={char_count}" + ) + logger.debug( + f"Table orientation rec scores table_index={table_index}: " + f"{'; '.join(score_parts)}" + ) + def _score_rotation_candidate_by_ocr(self, img_bgr: np.ndarray) -> tuple[float, int, int]: """对单个候选角度执行 OCR det+抽样 rec,返回平均置信度、有效文本数和字符数。""" task = self._build_orientation_score_task("", img_bgr) @@ -226,6 +245,7 @@ class MineruTableOrientationClsModel: rotated_img = self._rotate_image_by_label(img_bgr, label) score_by_label[label] = self._score_rotation_candidate_by_ocr(rotated_img) + self._debug_log_orientation_rec_scores(-1, score_by_label) return self._select_rotation_label_by_scores(score_by_label) @staticmethod @@ -401,6 +421,7 @@ class MineruTableOrientationClsModel: for img_info, tasks in img_score_tasks: score_by_label = self._score_orientation_tasks_with_rec(tasks, rec_res) + self._debug_log_orientation_rec_scores(img_info["index"], score_by_label) label_by_index[img_info["index"]] = self._select_rotation_label_by_scores( score_by_label ) From 8fecc190dcbaf4f247bde04756a7c61986d0dfaa Mon Sep 17 00:00:00 2001 From: myhloli Date: Wed, 17 Jun 2026 20:00:47 +0800 Subject: [PATCH 6/6] feat: refactor orientation scoring tasks and enhance batch detection logic --- .../model/table/cls/mineru_table_ori_cls.py | 191 +++++++++++++++--- 1 file changed, 158 insertions(+), 33 deletions(-) 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,