feat: implement table orientation classification and rotation handling

This commit is contained in:
myhloli
2026-05-13 00:32:26 +08:00
parent fc6b5892af
commit b17fcd0347
8 changed files with 427 additions and 351 deletions
+51 -11
View File
@@ -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"
)
# 表格分类
+6 -6
View File
@@ -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,
)
+1 -1
View File
@@ -7,4 +7,4 @@ class AtomicModel:
WirelessTable = "wireless_table"
WiredTable = "wired_table"
TableCls = "table_cls"
ImgOrientationCls = "img_ori_cls"
TableOrientationCls = "table_ori_cls"
-1
View File
@@ -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
View File
@@ -1 +0,0 @@
# Copyright (c) Opendatalab. All rights reserved.
-330
View File
@@ -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
-1
View File
@@ -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: