mirror of
https://github.com/opendatalab/MinerU.git
synced 2026-09-24 23:10:23 +08:00
feat: implement table orientation classification and rotation handling
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
import base64
|
||||
import html
|
||||
import re
|
||||
|
||||
import cv2
|
||||
from loguru import logger
|
||||
@@ -29,11 +30,15 @@ from ...utils.pdf_image_tools import get_crop_np_img
|
||||
LAYOUT_BASE_BATCH_SIZE = 1
|
||||
MFR_BASE_BATCH_SIZE = 16
|
||||
OCR_DET_BASE_BATCH_SIZE = 8
|
||||
TABLE_ORI_CLS_BATCH_SIZE = 16
|
||||
TABLE_Wired_Wireless_CLS_BATCH_SIZE = 16
|
||||
TABLE_OCR_REC_SINGLE_CHAR_REPLACEMENTS = {
|
||||
"香": "否",
|
||||
"哦樂": "哦",
|
||||
}
|
||||
TABLE_OCR_REC_REGEX_REPLACEMENTS = (
|
||||
# 仅规范化完整的“单个数字 + 號”,避免影响“10號”“第6號”等普通文本。
|
||||
(re.compile(r"^([0-9])號$"), r"\1"),
|
||||
)
|
||||
|
||||
|
||||
class BatchAnalyze:
|
||||
@@ -190,6 +195,28 @@ class BatchAnalyze:
|
||||
def _table_supports_inline_objects(table_res_dict: dict) -> bool:
|
||||
return str(table_res_dict.get("rotate_label", "0")) == "0"
|
||||
|
||||
@staticmethod
|
||||
def _apply_table_rotate_label(table_res_dict: dict, rotate_label: str) -> None:
|
||||
"""根据方向预测结果写回标签,并同步旋转无线和有线表格图片。"""
|
||||
rotate_label = str(rotate_label or "0")
|
||||
table_res_dict["rotate_label"] = rotate_label
|
||||
|
||||
if rotate_label == "270":
|
||||
rotate_code = cv2.ROTATE_90_CLOCKWISE
|
||||
elif rotate_label == "90":
|
||||
rotate_code = cv2.ROTATE_90_COUNTERCLOCKWISE
|
||||
else:
|
||||
return
|
||||
|
||||
table_res_dict["table_img"] = cv2.rotate(
|
||||
np.asarray(table_res_dict["table_img"]),
|
||||
rotate_code,
|
||||
)
|
||||
table_res_dict["wired_table_img"] = cv2.rotate(
|
||||
np.asarray(table_res_dict["wired_table_img"]),
|
||||
rotate_code,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _sort_table_ocr_result(ocr_result: list[list]) -> None:
|
||||
if not ocr_result:
|
||||
@@ -216,10 +243,16 @@ class BatchAnalyze:
|
||||
|
||||
@staticmethod
|
||||
def _normalize_table_ocr_rec_text(text):
|
||||
"""规范化表格 OCR rec 的已知单字误识别,避免后续表格模型消费错误文本。"""
|
||||
"""规范化表格 OCR rec 的已知误识别,避免后续表格模型消费错误文本。"""
|
||||
if not isinstance(text, str):
|
||||
return text
|
||||
return TABLE_OCR_REC_SINGLE_CHAR_REPLACEMENTS.get(text, text)
|
||||
if text in TABLE_OCR_REC_SINGLE_CHAR_REPLACEMENTS:
|
||||
return TABLE_OCR_REC_SINGLE_CHAR_REPLACEMENTS[text]
|
||||
for pattern, replacement in TABLE_OCR_REC_REGEX_REPLACEMENTS:
|
||||
match = pattern.fullmatch(text)
|
||||
if match:
|
||||
return match.expand(replacement)
|
||||
return text
|
||||
|
||||
@classmethod
|
||||
def _extract_table_inline_objects(
|
||||
@@ -426,21 +459,28 @@ class BatchAnalyze:
|
||||
if self.table_enable:
|
||||
|
||||
# 图片旋转批量处理
|
||||
img_orientation_cls_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.ImgOrientationCls,
|
||||
table_orientation_cls_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.TableOrientationCls,
|
||||
)
|
||||
try:
|
||||
if self.table_ori_cls_batch_enabled:
|
||||
img_orientation_cls_model.batch_predict(table_res_list_all_page,
|
||||
det_batch_size=self.batch_ratio * OCR_DET_BASE_BATCH_SIZE,
|
||||
batch_size=TABLE_ORI_CLS_BATCH_SIZE)
|
||||
rotate_labels = table_orientation_cls_model.batch_predict(
|
||||
table_res_list_all_page,
|
||||
det_batch_size=self.batch_ratio * OCR_DET_BASE_BATCH_SIZE,
|
||||
)
|
||||
if len(rotate_labels) != len(table_res_list_all_page):
|
||||
raise ValueError(
|
||||
"Table orientation batch prediction result count mismatch"
|
||||
)
|
||||
for table_res, rotate_label in zip(table_res_list_all_page, rotate_labels):
|
||||
self._apply_table_rotate_label(table_res, rotate_label)
|
||||
else:
|
||||
for table_res in table_res_list_all_page:
|
||||
rotate_label = img_orientation_cls_model.predict(table_res['table_img'])
|
||||
img_orientation_cls_model.img_rotate(table_res, rotate_label)
|
||||
rotate_label = table_orientation_cls_model.predict(table_res['table_img'])
|
||||
self._apply_table_rotate_label(table_res, rotate_label)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Image orientation classification failed: {e}, using original image"
|
||||
f"Table orientation classification failed: {e}, using original image"
|
||||
)
|
||||
|
||||
# 表格分类
|
||||
|
||||
@@ -10,8 +10,8 @@ from ...model.layout.pp_doclayoutv2 import PPDocLayoutV2LayoutModel
|
||||
from ...model.mfr.unimernet.Unimernet import UnimernetModel
|
||||
from ...model.mfr.pp_formulanet_plus_m.predict_formula import FormulaRecognizer
|
||||
from mineru.model.ocr.pytorch_paddle import PytorchPaddleOCR
|
||||
from ...model.ori_cls.paddle_ori_cls import PaddleOrientationClsModel
|
||||
from ...model.table.cls.paddle_table_cls import PaddleTableClsModel
|
||||
from ...model.table.cls.mineru_table_ori_cls import MineruTableOrientationClsModel
|
||||
from ...model.table.rec.slanet_plus.main import PaddleTableModel
|
||||
from ...model.table.rec.unet_table.main import UnetTableModel
|
||||
from ...utils.config_reader import get_device
|
||||
@@ -30,7 +30,7 @@ else:
|
||||
MFR_MODEL = "unimernet_small"
|
||||
|
||||
|
||||
def img_orientation_cls_model_init():
|
||||
def table_orientation_cls_model_init():
|
||||
atom_model_manager = AtomModelSingleton()
|
||||
ocr_engine = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.OCR,
|
||||
@@ -39,7 +39,7 @@ def img_orientation_cls_model_init():
|
||||
lang="ch_lite",
|
||||
enable_merge_det_boxes=False
|
||||
)
|
||||
cls_model = PaddleOrientationClsModel(ocr_engine)
|
||||
cls_model = MineruTableOrientationClsModel(ocr_engine)
|
||||
return cls_model
|
||||
|
||||
|
||||
@@ -183,8 +183,8 @@ def atom_model_init(model_name: str, **kwargs):
|
||||
)
|
||||
elif model_name == AtomicModel.TableCls:
|
||||
atom_model = table_cls_model_init()
|
||||
elif model_name == AtomicModel.ImgOrientationCls:
|
||||
atom_model = img_orientation_cls_model_init()
|
||||
elif model_name == AtomicModel.TableOrientationCls:
|
||||
atom_model = table_orientation_cls_model_init()
|
||||
else:
|
||||
logger.error('model name not allow')
|
||||
exit(1)
|
||||
@@ -253,7 +253,7 @@ class MineruPipelineModel:
|
||||
atom_model_name=AtomicModel.TableCls,
|
||||
)
|
||||
self.img_orientation_cls_model = atom_model_manager.get_atom_model(
|
||||
atom_model_name=AtomicModel.ImgOrientationCls,
|
||||
atom_model_name=AtomicModel.TableOrientationCls,
|
||||
lang=self.lang,
|
||||
)
|
||||
|
||||
|
||||
@@ -7,4 +7,4 @@ class AtomicModel:
|
||||
WirelessTable = "wireless_table"
|
||||
WiredTable = "wired_table"
|
||||
TableCls = "table_cls"
|
||||
ImgOrientationCls = "img_ori_cls"
|
||||
TableOrientationCls = "table_ori_cls"
|
||||
|
||||
@@ -72,7 +72,6 @@ def download_pipeline_models():
|
||||
ModelPath.slanet_plus,
|
||||
ModelPath.unet_structure,
|
||||
ModelPath.paddle_table_cls,
|
||||
ModelPath.paddle_orientation_classification,
|
||||
ModelPath.pp_formulanet_plus_m,
|
||||
]
|
||||
download_finish_path = ""
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
@@ -1,330 +0,0 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
import os
|
||||
|
||||
from PIL import Image
|
||||
from collections import defaultdict
|
||||
from typing import List, Dict
|
||||
from tqdm import tqdm
|
||||
import cv2
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
|
||||
from mineru.utils.enum_class import ModelPath
|
||||
from mineru.utils.models_download_utils import auto_download_and_get_model_root_path
|
||||
|
||||
|
||||
# 旋转门控只统计极窄 OCR 框,避免普通中文单字框被当作旋转证据。
|
||||
ROTATED_TEXT_ASPECT_RATIO_THRESHOLD = 0.58
|
||||
ROTATED_TEXT_RATIO_THRESHOLD = 0.015
|
||||
ROTATED_TEXT_MIN_BOXES = 5
|
||||
# 横向宽框占比足够高时认为是正常横向表格,直接豁免旋转检测。
|
||||
HORIZONTAL_TEXT_ASPECT_RATIO_THRESHOLD = 2.0
|
||||
HORIZONTAL_TEXT_REL_WIDTH_THRESHOLD = 0.06
|
||||
HORIZONTAL_TEXT_RATIO_THRESHOLD = 0.60
|
||||
|
||||
|
||||
class PaddleOrientationClsModel:
|
||||
def __init__(self, ocr_engine):
|
||||
self.sess = onnxruntime.InferenceSession(
|
||||
os.path.join(auto_download_and_get_model_root_path(ModelPath.paddle_orientation_classification), ModelPath.paddle_orientation_classification)
|
||||
)
|
||||
self.ocr_engine = ocr_engine
|
||||
self.less_length = 256
|
||||
self.cw, self.ch = 224, 224
|
||||
self.std = [0.229, 0.224, 0.225]
|
||||
self.scale = 0.00392156862745098
|
||||
self.mean = [0.485, 0.456, 0.406]
|
||||
self.labels = ["0", "90", "180", "270"]
|
||||
|
||||
def preprocess(self, input_img):
|
||||
# 放大图片,使其最短边长为256
|
||||
h, w = input_img.shape[:2]
|
||||
scale = 256 / min(h, w)
|
||||
h_resize = round(h * scale)
|
||||
w_resize = round(w * scale)
|
||||
img = cv2.resize(input_img, (w_resize, h_resize), interpolation=1)
|
||||
# 调整为224*224的正方形
|
||||
h, w = img.shape[:2]
|
||||
cw, ch = 224, 224
|
||||
x1 = max(0, (w - cw) // 2)
|
||||
y1 = max(0, (h - ch) // 2)
|
||||
x2 = min(w, x1 + cw)
|
||||
y2 = min(h, y1 + ch)
|
||||
if w < cw or h < ch:
|
||||
raise ValueError(
|
||||
f"Input image ({w}, {h}) smaller than the target size ({cw}, {ch})."
|
||||
)
|
||||
img = img[y1:y2, x1:x2, ...]
|
||||
# 正则化
|
||||
split_im = list(cv2.split(img))
|
||||
std = [0.229, 0.224, 0.225]
|
||||
scale = 0.00392156862745098
|
||||
mean = [0.485, 0.456, 0.406]
|
||||
alpha = [scale / std[i] for i in range(len(std))]
|
||||
beta = [-mean[i] / std[i] for i in range(len(std))]
|
||||
for c in range(img.shape[2]):
|
||||
split_im[c] = split_im[c].astype(np.float32)
|
||||
split_im[c] *= alpha[c]
|
||||
split_im[c] += beta[c]
|
||||
img = cv2.merge(split_im)
|
||||
# 5. 转换为 CHW 格式
|
||||
img = img.transpose((2, 0, 1))
|
||||
imgs = [img]
|
||||
x = np.stack(imgs, axis=0).astype(dtype=np.float32, copy=False)
|
||||
return x
|
||||
|
||||
def predict(self, input_img):
|
||||
rotate_label = "0" # Default to 0 if no rotation detected or not portrait
|
||||
if isinstance(input_img, Image.Image):
|
||||
np_img = np.asarray(input_img)
|
||||
elif isinstance(input_img, np.ndarray):
|
||||
np_img = input_img
|
||||
else:
|
||||
raise ValueError("Input must be a pillow object or a numpy array.")
|
||||
bgr_image = cv2.cvtColor(np_img, cv2.COLOR_RGB2BGR)
|
||||
# First check the overall image aspect ratio (height/width)
|
||||
img_height, img_width = bgr_image.shape[:2]
|
||||
img_aspect_ratio = img_height / img_width if img_width > 0 else 1.0
|
||||
img_is_portrait = img_aspect_ratio > 1.2
|
||||
|
||||
if img_is_portrait:
|
||||
|
||||
det_res = self.ocr_engine.ocr(bgr_image, rec=False)[0]
|
||||
# Check if table is rotated by analyzing text box aspect ratios
|
||||
if det_res:
|
||||
is_rotated = self._is_rotated_by_det_boxes(det_res, img_width)
|
||||
# logger.debug(f"Text orientation analysis: vertical={vertical_count}, det_res={len(det_res)}, rotated={is_rotated}")
|
||||
|
||||
# If we have more vertical text boxes than horizontal ones,
|
||||
# and vertical ones are significant, table might be rotated
|
||||
if is_rotated:
|
||||
x = self.preprocess(np_img)
|
||||
(result,) = self.sess.run(None, {"x": x})
|
||||
rotate_label = self._normalize_rotated_label(
|
||||
self.labels[np.argmax(result)]
|
||||
)
|
||||
# logger.debug(f"Orientation classification result: {label}")
|
||||
|
||||
return rotate_label
|
||||
|
||||
@staticmethod
|
||||
def _normalize_rotated_label(label: str) -> str:
|
||||
"""进入方向分类器已说明表格疑似旋转,0/180 等不可信结果统一按 270 兜底。"""
|
||||
if label in {"90", "270"}:
|
||||
return label
|
||||
return "270"
|
||||
|
||||
@staticmethod
|
||||
def _count_rotated_text_boxes(det_boxes) -> int:
|
||||
"""统计极窄 OCR 框数量,普通中文单字框不应被算作旋转表格证据。"""
|
||||
vertical_count = 0
|
||||
for box_ocr_res in det_boxes:
|
||||
p1, p2, p3, p4 = box_ocr_res
|
||||
|
||||
# 计算 OCR 框的宽高
|
||||
width = p3[0] - p1[0]
|
||||
height = p3[1] - p1[1]
|
||||
|
||||
aspect_ratio = width / height if height > 0 else 1.0
|
||||
|
||||
# 只统计足够极窄的竖向文字框
|
||||
if aspect_ratio < ROTATED_TEXT_ASPECT_RATIO_THRESHOLD:
|
||||
vertical_count += 1
|
||||
return vertical_count
|
||||
|
||||
@classmethod
|
||||
def _has_enough_horizontal_text_boxes(cls, det_boxes, image_width: int) -> bool:
|
||||
"""统计真实横向宽文本框占比,达标时豁免旋转检测。"""
|
||||
if image_width <= 0:
|
||||
return False
|
||||
|
||||
horizontal_count = 0
|
||||
valid_box_count = 0
|
||||
for box_ocr_res in det_boxes:
|
||||
p1, p2, p3, p4 = box_ocr_res
|
||||
|
||||
# 计算 OCR 框宽高,并跳过无效框,避免污染占比分母。
|
||||
width = p3[0] - p1[0]
|
||||
height = p3[1] - p1[1]
|
||||
if width <= 0 or height <= 0:
|
||||
continue
|
||||
|
||||
valid_box_count += 1
|
||||
aspect_ratio = width / height
|
||||
rel_width = width / image_width
|
||||
if (
|
||||
aspect_ratio >= HORIZONTAL_TEXT_ASPECT_RATIO_THRESHOLD
|
||||
and rel_width >= HORIZONTAL_TEXT_REL_WIDTH_THRESHOLD
|
||||
):
|
||||
horizontal_count += 1
|
||||
|
||||
return (
|
||||
valid_box_count > 0
|
||||
and horizontal_count >= valid_box_count * HORIZONTAL_TEXT_RATIO_THRESHOLD
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _is_rotated_by_det_boxes(cls, det_boxes, image_width: int) -> bool:
|
||||
"""用极窄 OCR 框的数量和占比判断是否进入方向分类器。"""
|
||||
if det_boxes is None or len(det_boxes) == 0:
|
||||
return False
|
||||
|
||||
if cls._has_enough_horizontal_text_boxes(det_boxes, image_width):
|
||||
return False
|
||||
|
||||
vertical_count = cls._count_rotated_text_boxes(det_boxes)
|
||||
return (
|
||||
vertical_count >= len(det_boxes) * ROTATED_TEXT_RATIO_THRESHOLD
|
||||
and vertical_count >= ROTATED_TEXT_MIN_BOXES
|
||||
)
|
||||
|
||||
def list_2_batch(self, img_list, batch_size=16):
|
||||
"""
|
||||
将任意长度的列表按照指定的batch size分成多个batch
|
||||
|
||||
Args:
|
||||
img_list: 输入的列表
|
||||
batch_size: 每个batch的大小,默认为16
|
||||
|
||||
Returns:
|
||||
一个包含多个batch的列表,每个batch都是原列表的一个子列表
|
||||
"""
|
||||
batches = []
|
||||
for i in range(0, len(img_list), batch_size):
|
||||
batch = img_list[i : min(i + batch_size, len(img_list))]
|
||||
batches.append(batch)
|
||||
return batches
|
||||
|
||||
def batch_preprocess(self, imgs):
|
||||
res_imgs = []
|
||||
for img_info in imgs:
|
||||
img = np.asarray(img_info["table_img"])
|
||||
# 放大图片,使其最短边长为256
|
||||
h, w = img.shape[:2]
|
||||
scale = 256 / min(h, w)
|
||||
h_resize = round(h * scale)
|
||||
w_resize = round(w * scale)
|
||||
img = cv2.resize(img, (w_resize, h_resize), interpolation=1)
|
||||
# 调整为224*224的正方形
|
||||
h, w = img.shape[:2]
|
||||
cw, ch = 224, 224
|
||||
x1 = max(0, (w - cw) // 2)
|
||||
y1 = max(0, (h - ch) // 2)
|
||||
x2 = min(w, x1 + cw)
|
||||
y2 = min(h, y1 + ch)
|
||||
if w < cw or h < ch:
|
||||
raise ValueError(
|
||||
f"Input image ({w}, {h}) smaller than the target size ({cw}, {ch})."
|
||||
)
|
||||
img = img[y1:y2, x1:x2, ...]
|
||||
# 正则化
|
||||
split_im = list(cv2.split(img))
|
||||
std = [0.229, 0.224, 0.225]
|
||||
scale = 0.00392156862745098
|
||||
mean = [0.485, 0.456, 0.406]
|
||||
alpha = [scale / std[i] for i in range(len(std))]
|
||||
beta = [-mean[i] / std[i] for i in range(len(std))]
|
||||
for c in range(img.shape[2]):
|
||||
split_im[c] = split_im[c].astype(np.float32)
|
||||
split_im[c] *= alpha[c]
|
||||
split_im[c] += beta[c]
|
||||
img = cv2.merge(split_im)
|
||||
# 5. 转换为 CHW 格式
|
||||
img = img.transpose((2, 0, 1))
|
||||
res_imgs.append(img)
|
||||
x = np.stack(res_imgs, axis=0).astype(dtype=np.float32, copy=False)
|
||||
return x
|
||||
|
||||
def batch_predict(
|
||||
self, imgs: List[Dict], det_batch_size: int, batch_size: int = 16
|
||||
) -> None:
|
||||
"""
|
||||
批量预测传入的包含图片信息列表的旋转信息,并且将旋转过的图片正确地旋转回来
|
||||
"""
|
||||
RESOLUTION_GROUP_STRIDE = 128
|
||||
# 跳过长宽比小于1.2的图片
|
||||
resolution_groups = defaultdict(list)
|
||||
for img in imgs:
|
||||
# RGB图像转换BGR
|
||||
bgr_img: np.ndarray = cv2.cvtColor(np.asarray(img["table_img"]), cv2.COLOR_RGB2BGR)
|
||||
img["table_img_bgr"] = bgr_img
|
||||
img_height, img_width = bgr_img.shape[:2]
|
||||
img_aspect_ratio = img_height / img_width if img_width > 0 else 1.0
|
||||
if img_aspect_ratio > 1.2:
|
||||
# 归一化尺寸到RESOLUTION_GROUP_STRIDE的倍数
|
||||
normalized_h = ((img_height + RESOLUTION_GROUP_STRIDE) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE # 向上取整到RESOLUTION_GROUP_STRIDE的倍数
|
||||
normalized_w = ((img_width + RESOLUTION_GROUP_STRIDE) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE
|
||||
group_key = (normalized_h, normalized_w)
|
||||
resolution_groups[group_key].append(img)
|
||||
|
||||
# 对每个分辨率组进行批处理
|
||||
rotated_imgs = []
|
||||
for group_key, group_imgs in tqdm(resolution_groups.items(), desc="Table-ori cls stage1 predict", disable=True):
|
||||
# 计算目标尺寸(组内最大尺寸,向上取整到RESOLUTION_GROUP_STRIDE的倍数)
|
||||
max_h = max(img["table_img_bgr"].shape[0] for img in group_imgs)
|
||||
max_w = max(img["table_img_bgr"].shape[1] for img in group_imgs)
|
||||
target_h = ((max_h + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE
|
||||
target_w = ((max_w + RESOLUTION_GROUP_STRIDE - 1) // RESOLUTION_GROUP_STRIDE) * RESOLUTION_GROUP_STRIDE
|
||||
|
||||
# 对所有图像进行padding到统一尺寸
|
||||
batch_images = []
|
||||
for img in group_imgs:
|
||||
bgr_img = img["table_img_bgr"]
|
||||
h, w = bgr_img.shape[:2]
|
||||
# 创建目标尺寸的白色背景
|
||||
padded_img = np.ones((target_h, target_w, 3), dtype=np.uint8) * 255
|
||||
# 将原图像粘贴到左上角
|
||||
padded_img[:h, :w] = bgr_img
|
||||
batch_images.append(padded_img)
|
||||
|
||||
# 批处理检测
|
||||
batch_results = self.ocr_engine.text_detector.batch_predict(
|
||||
batch_images, min(len(batch_images), det_batch_size)
|
||||
)
|
||||
|
||||
# 根据批处理结果检测图像是否旋转,旋转的图像放入列表中,继续进行旋转角度的预测
|
||||
|
||||
for index, (img_info, (dt_boxes, elapse)) in enumerate(
|
||||
zip(group_imgs, batch_results)
|
||||
):
|
||||
image_width = img_info["table_img_bgr"].shape[1]
|
||||
if self._is_rotated_by_det_boxes(dt_boxes, image_width):
|
||||
rotated_imgs.append(img_info)
|
||||
|
||||
# 对旋转的图片进行旋转角度预测
|
||||
if len(rotated_imgs) > 0:
|
||||
imgs = self.list_2_batch(rotated_imgs, batch_size=batch_size)
|
||||
with tqdm(total=len(rotated_imgs), desc="Table-ori cls stage2 predict", disable=True) as pbar:
|
||||
for img_batch in imgs:
|
||||
x = self.batch_preprocess(img_batch)
|
||||
results = self.sess.run(None, {"x": x})
|
||||
for img_info, res in zip(rotated_imgs, results[0]):
|
||||
label = self._normalize_rotated_label(
|
||||
self.labels[np.argmax(res)]
|
||||
)
|
||||
self.img_rotate(img_info, label)
|
||||
pbar.update(1)
|
||||
|
||||
def img_rotate(self, img_info, label):
|
||||
img_info["rotate_label"] = label
|
||||
if label == "270":
|
||||
img_info["table_img"] = cv2.rotate(
|
||||
np.asarray(img_info["table_img"]),
|
||||
cv2.ROTATE_90_CLOCKWISE,
|
||||
)
|
||||
img_info["wired_table_img"] = cv2.rotate(
|
||||
np.asarray(img_info["wired_table_img"]),
|
||||
cv2.ROTATE_90_CLOCKWISE,
|
||||
)
|
||||
elif label == "90":
|
||||
img_info["table_img"] = cv2.rotate(
|
||||
np.asarray(img_info["table_img"]),
|
||||
cv2.ROTATE_90_COUNTERCLOCKWISE,
|
||||
)
|
||||
img_info["wired_table_img"] = cv2.rotate(
|
||||
np.asarray(img_info["wired_table_img"]),
|
||||
cv2.ROTATE_90_COUNTERCLOCKWISE,
|
||||
)
|
||||
else:
|
||||
# 180度和0度不做处理
|
||||
pass
|
||||
@@ -0,0 +1,369 @@
|
||||
# Copyright (c) Opendatalab. All rights reserved.
|
||||
|
||||
from PIL import Image
|
||||
from collections import defaultdict
|
||||
from typing import List, Dict
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
# 旋转候选门控回到旧规则,先尽量召回疑似旋转表,再由 OCR rec 评分决定最终角度。
|
||||
ROTATED_TEXT_ASPECT_RATIO_THRESHOLD = 0.8
|
||||
ROTATED_TEXT_RATIO_THRESHOLD = 0.28
|
||||
ROTATED_TEXT_MIN_BOXES = 3
|
||||
# OCR rec 角度评分参数,控制抽样成本和 0 度优先的保守阈值。
|
||||
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_LABELS = ("0", "90", "270")
|
||||
|
||||
|
||||
class MineruTableOrientationClsModel:
|
||||
def __init__(self, ocr_engine):
|
||||
self.ocr_engine = ocr_engine
|
||||
|
||||
def predict(self, input_img):
|
||||
np_img = self._to_numpy_image(input_img)
|
||||
|
||||
# 单张预测作为 batch_predict 的特例,保证门控、det 和 OCR 评分逻辑完全一致。
|
||||
return self.batch_predict([{"table_img": np_img}], det_batch_size=1)[0]
|
||||
|
||||
@staticmethod
|
||||
def _to_numpy_image(input_img) -> np.ndarray:
|
||||
"""统一将 Pillow/ndarray 输入转为 numpy 图像,保持外部入参校验一致。"""
|
||||
if isinstance(input_img, Image.Image):
|
||||
return np.asarray(input_img)
|
||||
if isinstance(input_img, np.ndarray):
|
||||
return input_img
|
||||
raise ValueError("Input must be a pillow object or a numpy array.")
|
||||
|
||||
@classmethod
|
||||
def _to_bgr_table_image(cls, table_info: Dict) -> np.ndarray:
|
||||
"""从表格信息中读取 table_img,并转换为 OCR detector 使用的 BGR 图像。"""
|
||||
table_img = cls._to_numpy_image(table_info["table_img"])
|
||||
return cv2.cvtColor(table_img, cv2.COLOR_RGB2BGR)
|
||||
|
||||
@staticmethod
|
||||
def _ceil_to_stride(value: int, stride: int) -> int:
|
||||
"""将尺寸向上对齐到 stride 倍数,已经整除时保持原尺寸。"""
|
||||
if stride <= 0:
|
||||
raise ValueError("stride must be positive")
|
||||
if value <= 0:
|
||||
return 0
|
||||
return ((value + stride - 1) // stride) * stride
|
||||
|
||||
@staticmethod
|
||||
def _box_width_height(box_ocr_res) -> tuple[float, float]:
|
||||
"""从 OCR 四点框中提取宽高,统一处理 list/ndarray 两种输入。"""
|
||||
points = np.asarray(box_ocr_res, dtype=np.float32)
|
||||
p1 = points[0]
|
||||
p3 = points[2]
|
||||
return float(p3[0] - p1[0]), float(p3[1] - p1[1])
|
||||
|
||||
@staticmethod
|
||||
def _count_rotated_text_boxes(det_boxes) -> int:
|
||||
"""统计符合旧规则的高窄 OCR 框数量,作为疑似旋转表候选证据。"""
|
||||
vertical_count = 0
|
||||
for box_ocr_res in det_boxes:
|
||||
width, height = MineruTableOrientationClsModel._box_width_height(box_ocr_res)
|
||||
aspect_ratio = width / height if height > 0 else 1.0
|
||||
|
||||
# 旧规则允许更宽的高窄框进入候选,最终是否旋转交给 OCR rec 评分。
|
||||
if aspect_ratio < ROTATED_TEXT_ASPECT_RATIO_THRESHOLD:
|
||||
vertical_count += 1
|
||||
return vertical_count
|
||||
|
||||
@classmethod
|
||||
def _is_rotation_candidate_by_det_boxes(cls, det_boxes) -> bool:
|
||||
"""用旧竖框规则判断是否进入 OCR 多角度评分。"""
|
||||
if det_boxes is None or len(det_boxes) == 0:
|
||||
return False
|
||||
|
||||
vertical_count = cls._count_rotated_text_boxes(det_boxes)
|
||||
return (
|
||||
vertical_count >= len(det_boxes) * ROTATED_TEXT_RATIO_THRESHOLD
|
||||
and vertical_count >= ROTATED_TEXT_MIN_BOXES
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _rotate_image_by_label(img: np.ndarray, label: str) -> np.ndarray:
|
||||
"""按候选角度旋转图像,0 度返回副本以避免后续误改原图。"""
|
||||
if label == "270":
|
||||
return cv2.rotate(img, cv2.ROTATE_90_CLOCKWISE)
|
||||
if label == "90":
|
||||
return cv2.rotate(img, cv2.ROTATE_90_COUNTERCLOCKWISE)
|
||||
return img.copy()
|
||||
|
||||
@staticmethod
|
||||
def _sample_det_boxes(det_boxes) -> list:
|
||||
"""对 OCR det 框做均匀抽样,限制 rec 打分成本并覆盖整张表。"""
|
||||
if det_boxes is None or len(det_boxes) == 0:
|
||||
return []
|
||||
if len(det_boxes) <= ORIENTATION_SCORE_MAX_SAMPLE_BOXES:
|
||||
return list(det_boxes)
|
||||
|
||||
indexes = np.linspace(
|
||||
0,
|
||||
len(det_boxes) - 1,
|
||||
ORIENTATION_SCORE_MAX_SAMPLE_BOXES,
|
||||
)
|
||||
sampled_indexes = sorted({int(round(index)) for index in indexes})
|
||||
return [det_boxes[index] for index in sampled_indexes]
|
||||
|
||||
@staticmethod
|
||||
def _crop_image_without_text_rotation(img: np.ndarray, points) -> np.ndarray | None:
|
||||
"""为方向评分按外接矩形切图,不做透视修正或文本方向自动转正。"""
|
||||
points = np.asarray(points, dtype=np.float32)
|
||||
if len(points) != 4:
|
||||
return None
|
||||
|
||||
img_height, img_width = img.shape[:2]
|
||||
xmin = max(0, int(np.floor(np.min(points[:, 0]))))
|
||||
xmax = min(img_width, int(np.ceil(np.max(points[:, 0]))))
|
||||
ymin = max(0, int(np.floor(np.min(points[:, 1]))))
|
||||
ymax = min(img_height, int(np.ceil(np.max(points[:, 1]))))
|
||||
if xmax <= xmin or ymax <= ymin:
|
||||
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)
|
||||
|
||||
img_crop_list = []
|
||||
for box in sampled_boxes:
|
||||
crop_img = self._crop_image_without_text_rotation(img_bgr, box)
|
||||
if crop_img is not None and crop_img.size > 0:
|
||||
img_crop_list.append(crop_img)
|
||||
|
||||
return {
|
||||
"label": label,
|
||||
"crops": img_crop_list,
|
||||
"crop_count": len(img_crop_list),
|
||||
"crop_start": 0,
|
||||
"crop_end": len(img_crop_list),
|
||||
}
|
||||
|
||||
def _build_orientation_score_tasks(self, img_bgr: np.ndarray) -> List[Dict]:
|
||||
"""为一张表构造 0/90/270 三个角度的评分任务。"""
|
||||
tasks = []
|
||||
for label in ORIENTATION_SCORE_LABELS:
|
||||
rotated_img = self._rotate_image_by_label(img_bgr, label)
|
||||
tasks.append(self._build_orientation_score_task(label, rotated_img))
|
||||
return tasks
|
||||
|
||||
@staticmethod
|
||||
def _score_rec_results(rec_res) -> tuple[float, int, int]:
|
||||
"""根据 OCR rec 结果计算平均置信度、有效文本数和字符数。"""
|
||||
valid_scores = []
|
||||
char_count = 0
|
||||
for rec_item in rec_res or []:
|
||||
if not rec_item or len(rec_item) < 2:
|
||||
continue
|
||||
text, score = rec_item
|
||||
text = str(text)
|
||||
if not text.strip():
|
||||
continue
|
||||
valid_scores.append(float(score))
|
||||
char_count += len(text)
|
||||
|
||||
if len(valid_scores) < ORIENTATION_SCORE_MIN_VALID_RESULTS:
|
||||
return 0.0, len(valid_scores), char_count
|
||||
|
||||
return float(np.mean(valid_scores)), len(valid_scores), char_count
|
||||
|
||||
def _score_orientation_tasks_with_rec(self, tasks: List[Dict], rec_res) -> Dict[str, tuple[float, int, int]]:
|
||||
"""按任务记录的 crop slice 回填 rec 结果,得到每个角度的评分。"""
|
||||
score_by_label = {}
|
||||
rec_res = rec_res or []
|
||||
for task in tasks:
|
||||
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]
|
||||
score_by_label[task["label"]] = self._score_rec_results(task_rec_res)
|
||||
return score_by_label
|
||||
|
||||
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)
|
||||
if task["crop_count"] == 0:
|
||||
return 0.0, 0, 0
|
||||
|
||||
rec_ocr_res = self.ocr_engine.ocr(task["crops"], det=False, rec=True)
|
||||
rec_res = rec_ocr_res[0] if rec_ocr_res else None
|
||||
return self._score_rec_results(rec_res)
|
||||
|
||||
@staticmethod
|
||||
def _select_rotation_label_by_scores(score_by_label: Dict[str, tuple[float, int, int]]) -> str:
|
||||
"""按 OCR 评分选择最终角度,分差较小时优先保持 0 度。"""
|
||||
if not score_by_label:
|
||||
return "0"
|
||||
|
||||
zero_score = score_by_label.get("0", (0.0, 0, 0))[0]
|
||||
if zero_score >= ORIENTATION_ZERO_SCORE_PRIORITY_THRESHOLD:
|
||||
return "0"
|
||||
|
||||
best_label = max(
|
||||
ORIENTATION_SCORE_LABELS,
|
||||
key=lambda label: score_by_label.get(label, (0.0, 0, 0)),
|
||||
)
|
||||
best_score = score_by_label.get(best_label, (0.0, 0, 0))[0]
|
||||
if (
|
||||
best_label != "0"
|
||||
and best_score - zero_score < ORIENTATION_SCORE_TIE_THRESHOLD
|
||||
):
|
||||
return "0"
|
||||
return best_label
|
||||
|
||||
def _select_rotation_by_ocr_score(self, img_bgr: np.ndarray) -> str:
|
||||
"""比较 0/90/270 三个角度的 OCR rec 分数,分差很小时优先保持 0 度。"""
|
||||
score_by_label = {}
|
||||
for label in ORIENTATION_SCORE_LABELS:
|
||||
rotated_img = self._rotate_image_by_label(img_bgr, label)
|
||||
score_by_label[label] = self._score_rotation_candidate_by_ocr(rotated_img)
|
||||
|
||||
return self._select_rotation_label_by_scores(score_by_label)
|
||||
|
||||
@classmethod
|
||||
def _collect_portrait_image_groups(
|
||||
cls,
|
||||
imgs: List[Dict],
|
||||
resolution_group_stride: int,
|
||||
) -> Dict[tuple[int, int], list[Dict]]:
|
||||
"""按归一化分辨率收集竖版表格,横版表格默认保持 0 度跳过后续 OCR。"""
|
||||
resolution_groups = defaultdict(list)
|
||||
for index, img in enumerate(imgs):
|
||||
bgr_img = cls._to_bgr_table_image(img)
|
||||
img_height, img_width = bgr_img.shape[:2]
|
||||
img_aspect_ratio = img_height / img_width if img_width > 0 else 1.0
|
||||
if img_aspect_ratio <= 1.2:
|
||||
continue
|
||||
|
||||
group_key = (
|
||||
cls._ceil_to_stride(img_height, resolution_group_stride),
|
||||
cls._ceil_to_stride(img_width, resolution_group_stride),
|
||||
)
|
||||
resolution_groups[group_key].append(
|
||||
{
|
||||
"index": index,
|
||||
"table_img_bgr": bgr_img,
|
||||
}
|
||||
)
|
||||
return resolution_groups
|
||||
|
||||
@classmethod
|
||||
def _pad_group_images(
|
||||
cls,
|
||||
group_imgs: list[Dict],
|
||||
resolution_group_stride: int,
|
||||
) -> list[np.ndarray]:
|
||||
"""将同组表格 padding 到统一尺寸,便于 OCR detector 批处理。"""
|
||||
max_h = max(img["table_img_bgr"].shape[0] for img in group_imgs)
|
||||
max_w = max(img["table_img_bgr"].shape[1] for img in group_imgs)
|
||||
target_h = cls._ceil_to_stride(max_h, resolution_group_stride)
|
||||
target_w = cls._ceil_to_stride(max_w, resolution_group_stride)
|
||||
|
||||
batch_images = []
|
||||
for img in group_imgs:
|
||||
bgr_img = img["table_img_bgr"]
|
||||
h, w = bgr_img.shape[:2]
|
||||
padded_img = np.ones((target_h, target_w, 3), dtype=np.uint8) * 255
|
||||
padded_img[:h, :w] = bgr_img
|
||||
batch_images.append(padded_img)
|
||||
return batch_images
|
||||
|
||||
def _detect_rotation_candidates(
|
||||
self,
|
||||
resolution_groups: Dict[tuple[int, int], list[Dict]],
|
||||
det_batch_size: int,
|
||||
resolution_group_stride: int,
|
||||
) -> 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)),
|
||||
)
|
||||
|
||||
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)
|
||||
return rotated_imgs
|
||||
|
||||
def _build_score_tasks_for_candidates(
|
||||
self,
|
||||
rotated_imgs: list[Dict],
|
||||
) -> tuple[list[tuple[Dict, list[Dict]]], list[np.ndarray]]:
|
||||
"""为所有旋转候选构造三角度评分任务,并汇总成一次 OCR rec 输入。"""
|
||||
img_score_tasks = []
|
||||
all_crop_imgs = []
|
||||
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
|
||||
img_score_tasks.append((img_info, tasks))
|
||||
return img_score_tasks, all_crop_imgs
|
||||
|
||||
def _recognize_orientation_crops(self, all_crop_imgs: list[np.ndarray]) -> list:
|
||||
"""对所有候选角度 crop 合并执行 OCR rec,返回可按 slice 回填的结果。"""
|
||||
if not all_crop_imgs:
|
||||
return []
|
||||
|
||||
rec_ocr_res = self.ocr_engine.ocr(
|
||||
all_crop_imgs,
|
||||
det=False,
|
||||
rec=True,
|
||||
)
|
||||
return rec_ocr_res[0] if rec_ocr_res else []
|
||||
|
||||
def _score_rotation_candidates(self, rotated_imgs: list[Dict]) -> Dict[int, str]:
|
||||
"""批量评分旋转候选,并返回原始表格下标到最终角度标签的映射。"""
|
||||
if not rotated_imgs:
|
||||
return {}
|
||||
|
||||
label_by_index = {}
|
||||
img_score_tasks, all_crop_imgs = self._build_score_tasks_for_candidates(
|
||||
rotated_imgs
|
||||
)
|
||||
rec_res = self._recognize_orientation_crops(all_crop_imgs)
|
||||
|
||||
for img_info, tasks in img_score_tasks:
|
||||
score_by_label = self._score_orientation_tasks_with_rec(tasks, rec_res)
|
||||
label_by_index[img_info["index"]] = self._select_rotation_label_by_scores(
|
||||
score_by_label
|
||||
)
|
||||
return label_by_index
|
||||
|
||||
def batch_predict(
|
||||
self,
|
||||
imgs: List[Dict],
|
||||
det_batch_size: int,
|
||||
) -> List[str]:
|
||||
"""
|
||||
批量预测传入表格图片的旋转角度,只返回角度,不修改输入图片。
|
||||
"""
|
||||
RESOLUTION_GROUP_STRIDE = 128
|
||||
rotate_labels = ["0"] * len(imgs)
|
||||
resolution_groups = self._collect_portrait_image_groups(
|
||||
imgs,
|
||||
RESOLUTION_GROUP_STRIDE,
|
||||
)
|
||||
rotated_imgs = self._detect_rotation_candidates(
|
||||
resolution_groups,
|
||||
det_batch_size,
|
||||
RESOLUTION_GROUP_STRIDE,
|
||||
)
|
||||
label_by_index = self._score_rotation_candidates(rotated_imgs)
|
||||
for index, label in label_by_index.items():
|
||||
rotate_labels[index] = label
|
||||
|
||||
return rotate_labels
|
||||
@@ -108,7 +108,6 @@ class ModelPath:
|
||||
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"
|
||||
paddle_orientation_classification = "models/OriCls/paddle_orientation_classification/PP-LCNet_x1_0_doc_ori.onnx"
|
||||
|
||||
|
||||
class SplitFlag:
|
||||
|
||||
Reference in New Issue
Block a user