Files
2026-08-31 16:41:12 +08:00

527 lines
24 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
OCR引擎模块
使用 RapidOCR 对整张图片进行文字识别。
推理设备支持:CPU / CUDA / DirectML / TensorRT
"""
import numpy as np
import cv2
from rapidocr import EngineType, LangCls, LangDet, LangRec, ModelType, OCRVersion, RapidOCR
from rapidocr.utils.output import RapidOCROutput
from pathlib import Path
import logging
logger = logging.getLogger(__name__)
class OCRProcessor:
"""整图 OCR 处理器"""
def __init__(
self,
use_det: bool = True,
use_cls: bool = False,
use_rec: bool = True,
model_dir: str | Path = "model",
text_score: float = 0.5,
max_side_len: int = 2000,
det_limit_side_len: int = 736,
det_box_thresh: float = 0.5,
det_thresh: float = 0.3,
rec_batch_num: int = 6,
# 推理设备
device: str = "cpu",
device_id: int = 0,
intra_threads: int = 4,
inter_threads: int = 2,
# 竖排文字补充通道
enable_vertical: bool = False,
):
self.use_det = use_det
self.use_cls = use_cls
self.use_rec = use_rec
self.device = device
self.device_id = device_id
self.intra_threads = intra_threads
self.inter_threads = inter_threads
self.text_score = text_score
self.max_side_len = max_side_len
self.det_limit_side_len = det_limit_side_len
self.det_box_thresh = det_box_thresh
self.det_thresh = det_thresh
self.rec_batch_num = rec_batch_num
self.enable_vertical = enable_vertical
self._vert_engine = None
# 模型路径
model_path = Path(model_dir).resolve()
self.rec_model_path = model_path / "ch_PP-OCRv5_rec_mobile.onnx"
self.det_model_path = model_path / "ch_PP-OCRv5_det_mobile.onnx"
self.cls_model_path = model_path / "ch_PP-LCNet_x0_25_textline_ori_cls_mobile.onnx"
missing_models = [
str(path)
for path in (self.det_model_path, self.cls_model_path, self.rec_model_path)
if not path.is_file()
]
if missing_models:
raise FileNotFoundError(f"OCR模型文件不存在: {', '.join(missing_models)}")
# 初始化引擎
self.engine = self._init_engine()
def _build_device_params(self) -> dict:
"""根据 device 配置构建推理引擎参数"""
params = {}
if self.device == "cpu":
# CPU: 控制线程数以优化并发
params.update({
"EngineConfig.onnxruntime.intra_op_num_threads": self.intra_threads,
"EngineConfig.onnxruntime.inter_op_num_threads": self.inter_threads,
})
logger.info(f"推理设备: CPU (intra_threads={self.intra_threads}, inter_threads={self.inter_threads})")
elif self.device == "cuda":
# NVIDIA GPU with CUDA
params.update({
"EngineConfig.onnxruntime.use_cuda": True,
"EngineConfig.onnxruntime.cuda_ep_cfg.device_id": self.device_id,
})
logger.info(f"推理设备: CUDA (device_id={self.device_id})")
elif self.device == "dml":
# Windows DirectML (任意GPU)
params.update({
"EngineConfig.onnxruntime.use_dml": True,
})
logger.info("推理设备: DirectML")
elif self.device == "tensorrt":
# NVIDIA TensorRT (需要预先构建engine)
params.update({
"Det.engine_type": EngineType.TENSORRT,
"Rec.engine_type": EngineType.TENSORRT,
"Cls.engine_type": EngineType.TENSORRT,
"EngineConfig.tensorrt.device_id": self.device_id,
"EngineConfig.tensorrt.use_fp16": True,
})
logger.info(f"推理设备: TensorRT (device_id={self.device_id}, fp16=True)")
return params
def _init_engine(self) -> RapidOCR:
"""初始化OCR引擎"""
params = {
# 全局与常用调优参数
"Global.text_score": self.text_score,
"Global.max_side_len": self.max_side_len,
"Det.limit_side_len": self.det_limit_side_len,
"Det.box_thresh": self.det_box_thresh,
"Det.thresh": self.det_thresh,
"Rec.rec_batch_num": self.rec_batch_num,
# 模型契约
"Det.model_path": str(self.det_model_path),
"Det.engine_type": EngineType.ONNXRUNTIME,
"Det.lang_type": LangDet.CH,
"Det.model_type": ModelType.MOBILE,
"Det.ocr_version": OCRVersion.PPOCRV5,
"Rec.model_path": str(self.rec_model_path),
"Rec.engine_type": EngineType.ONNXRUNTIME,
"Rec.lang_type": LangRec.CH,
"Rec.model_type": ModelType.MOBILE,
"Rec.ocr_version": OCRVersion.PPOCRV5,
"Cls.engine_type": EngineType.ONNXRUNTIME,
"Cls.lang_type": LangCls.CH,
"Cls.model_type": ModelType.MOBILE,
"Cls.ocr_version": OCRVersion.PPOCRV5,
"Cls.model_path": str(self.cls_model_path),
}
# 合并设备参数
params.update(self._build_device_params())
return RapidOCR(params=params)
def process(self, image: np.ndarray) -> RapidOCROutput:
"""标准OCR识别(整图直接OCR)"""
return self.engine(
image,
use_det=self.use_det,
use_rec=self.use_rec,
use_cls=self.use_cls,
return_word_box=True
)
# ─── 边缘/竖排文字补充通道(纯 CV,CPU 友好)──────────
@staticmethod
def _row_bands(gray_img: np.ndarray, dark_thresh: int = 160, char_h: int | None = None):
"""行投影找文字行带(字高自适应)。
手写/毛笔大字内部笔画断裂会产生大间隙,固定空白阈值会把一个字拆成多段。
流程:小阈值粗切 → 按字高比例合并相邻带、丢弃碎片带。
char_h 为字高(px);None 时自动估计(取主带——≥0.35×最大带高——的平均高,
避免被小噪声碎片带拉偏)。
返回 [(y0, y1), ...](区域内坐标)
"""
row_dark = (gray_img < dark_thresh).sum(axis=1)
bands = []
in_b = False
for y, d in enumerate(row_dark):
if d > 2 and not in_b:
in_b = True
b0 = y
elif d <= 2 and in_b:
in_b = False
if y - b0 >= 2:
bands.append((b0, y))
if in_b:
bands.append((b0, len(row_dark)))
if not bands:
return []
if char_h is None:
hmax = max(y1 - y0 for y0, y1 in bands)
mains = [y1 - y0 for y0, y1 in bands if y1 - y0 >= hmax * 0.35]
char_h = round(sum(mains) / len(mains)) if mains else hmax
gap_min = max(2, int(char_h * 0.35)) # 相邻带空白 <= 该值 -> 笔画断裂,合并
band_min = max(4, int(char_h * 0.5)) # 低于该高度的带为噪声碎片,丢弃
merged = []
last = None
for y0, y1 in bands:
if y1 - y0 < band_min:
continue
if last is not None and y0 - last[1] <= gap_min:
last[1] = y1
else:
last = [y0, y1]
merged.append(last)
return [(y0, y1) for y0, y1 in merged]
def _cv_supplement(self, image: np.ndarray, main_boxes: list, band_h: int = 160) -> list:
"""边缘带补充检测(纯 OpenCV,无模型依赖,CPU 毫秒级):
四边带 → 多级低阈值二值化 → 膨胀 → 连通域 → 文字状块过滤。
解决主流程漏检的边缘小字/竖排手写文字。
返回 [(box_np, text), ...](原图坐标)"""
if not self.enable_vertical:
return []
H, W = image.shape[:2]
gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
bands = [
(0, 0, W, min(band_h, H)), # 上边带
(0, max(0, H - band_h), W, H), # 下边带
(0, 0, min(band_h, W), H), # 左边带
(max(0, W - band_h), 0, W, H), # 右边带
]
extra = []
for bx1, by1, bx2, by2 in bands:
if bx2 - bx1 < 40 or by2 - by1 < 40:
continue
band = gray[by1:by2, bx1:bx2]
for th in (155, 145, 135):
_, bw = cv2.threshold(band, th, 255, cv2.THRESH_BINARY_INV)
kernel = cv2.getStructuringElement(cv2.MORPH_RECT, (3, 5))
closed = cv2.morphologyEx(bw, cv2.MORPH_CLOSE, kernel, iterations=2)
contours, _ = cv2.findContours(closed, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
for c in contours:
x, y, w, h = cv2.boundingRect(c)
if not (8 <= w <= 110 and 10 <= h <= 130):
continue
if closed[y:y + h, x:x + w].mean() / 255.0 < 0.12:
continue
cx, cy = bx1 + x + w / 2, by1 + y + h / 2
# 块中心落在任一主流程检测框内(含4px容差)则跳过
if any(b[:, 0].min() - 4 <= cx <= b[:, 0].max() + 4
and b[:, 1].min() - 4 <= cy <= b[:, 1].max() + 4 for b in main_boxes):
continue
if any(abs(cx - eb[:, 0].mean()) < 10 and abs(cy - eb[:, 1].mean()) < 10 for eb, _ in extra):
continue
box = np.array([[x, y], [x + w, y], [x + w, y + h], [x, y + h]], dtype=float)
box[:, 0] += bx1
box[:, 1] += by1
text = self._cv_box_text(gray, box) if h / max(w, 1) > 1.2 else ""
extra.append((box, text))
return extra
def _cv_box_text(self, gray: np.ndarray, box: np.ndarray) -> str:
"""补充框内逐字识别(mobile rec,CPU 可跑):行投影切单字 → 逐字识别。
识别失败返回空串(标注优先,文本对错次要)。"""
H, W = gray.shape[:2]
x1, y1 = int(max(0, box[:, 0].min() - 3)), int(max(0, box[:, 1].min() - 3))
x2, y2 = int(min(W, box[:, 0].max() + 3)), int(min(H, box[:, 1].max() + 3))
if x2 - x1 < 4 or y2 - y1 < 4:
return ""
region = gray[y1:y2, x1:x2]
text = ""
for c0, c1 in self._row_bands(region, 160):
seg = region[c0:c1, :]
if seg.shape[0] < 8:
continue
big = cv2.resize(seg, None, fx=2, fy=2, interpolation=cv2.INTER_LANCZOS4)
r = self.engine(cv2.cvtColor(big, cv2.COLOR_GRAY2RGB), use_det=False, use_cls=False, use_rec=True)
if r.txts:
text += "".join(r.txts)
return text
def _vertical_box_text(self, gray: np.ndarray, box: np.ndarray) -> str:
"""竖排框内逐字识别:行投影(阈值160)切单字 → 逐字 mobile rec"""
H, W = gray.shape[:2]
x1, y1 = int(max(0, box[:, 0].min() - 3)), int(max(0, box[:, 1].min() - 3))
x2, y2 = int(min(W, box[:, 0].max() + 3)), int(min(H, box[:, 1].max() + 3))
if x2 - x1 < 4 or y2 - y1 < 4:
return ""
region = gray[y1:y2, x1:x2]
text = ""
for c0, c1 in self._row_bands(region, 160):
seg = region[c0:c1, :]
if seg.shape[0] < 8:
continue
big = cv2.resize(seg, None, fx=2, fy=2, interpolation=cv2.INTER_LANCZOS4)
r = self.engine(cv2.cvtColor(big, cv2.COLOR_GRAY2RGB), use_det=False, use_cls=False, use_rec=True)
if r.txts:
text += "".join(r.txts)
return text
def _vertical_longbox_supplement(self, image: np.ndarray, result: RapidOCROutput) -> RapidOCROutput:
"""竖排长框旋转识别补充:
检测把竖排一列文字框成超长行(>300px),Rec 压缩后丢字(如"殁")。
按列合并长框为竖条 → 旋转90°变横排 → 完整OCR → 行框换算回原图 →
行内均分字坐标 → 补充原 word_results 缺失的位置。"""
if result.boxes is None or not result.word_results:
return result
H, W = image.shape[:2]
# 超长竖排框(>300px 高且高宽比 >2)
def is_longbox(b):
return (b[:, 1].max() - b[:, 1].min() > 300
and (b[:, 1].max() - b[:, 1].min()) / max(b[:, 0].max() - b[:, 0].min(), 1) > 2)
def flat_items(word_results):
"""展平 word_results;rapidocr 对超长竖框的字级切分会把整列切成
碎片(半字高),这些框的字级条目整框丢弃,交给下方长框通道重切"""
out = []
for i, line in enumerate(word_results):
if isinstance(line, (tuple, list)) and len(line) >= 3 and isinstance(line[0], str):
out.append((line[0], line[1], line[2]))
elif isinstance(line, (tuple, list)):
if result.boxes is not None and i < len(result.boxes) and is_longbox(result.boxes[i]):
continue
for it in line:
if isinstance(it, (tuple, list)) and len(it) >= 3 and isinstance(it[0], str):
out.append((it[0], it[1], it[2]))
return out
def center(coords):
return sum(p[0] for p in coords) / len(coords), sum(p[1] for p in coords) / len(coords)
orig_items = flat_items(result.word_results)
# 超长竖排框(>300px 高且高宽比 >2)按 x 中心聚类合并为竖条
longs = [b for b in result.boxes if is_longbox(b)]
if not longs:
return result
# 聚类容差按列宽自适应: 硬编码 60px 会把列间距较密的竖排相邻列误并为一组
col_w = float(np.median([b[:, 0].max() - b[:, 0].min() for b in longs]))
cluster_d = max(2, int(col_w * 0.8))
groups = []
for b in longs:
cx = b[:, 0].mean()
for g in groups:
if abs(cx - g[0]) < cluster_d:
g[1] = min(g[1], b[:, 1].min()); g[2] = max(g[2], b[:, 1].max())
g[3] = min(g[3], b[:, 0].min()); g[4] = max(g[4], b[:, 0].max())
break
else:
groups.append([cx, b[:, 1].min(), b[:, 1].max(), b[:, 0].min(), b[:, 0].max()])
for cx, gy1, gy2, gx1, gx2 in groups:
x1, y1 = max(0, int(gx1)), max(0, int(gy1))
x2, y2 = min(W, int(gx2)), min(H, int(gy2))
if x2 - x1 < 10 or y2 - y1 < 50:
continue
crop = image[y1:y2, x1:x2]
_, chw = crop.shape[:2]
# 竖排(上→下)需逆时针转90°: 顺时针会字头朝左、字序反转, rec 无法识别
r90 = cv2.rotate(crop, cv2.ROTATE_90_COUNTERCLOCKWISE)
big = cv2.resize(r90, None, fx=2, fy=2, interpolation=cv2.INTER_CUBIC)
# 旋转后为横排: 列投影切字(按字高比例合并断笔) -> 逐字识别。
# 整列 rec + 均分会把毛笔大字切成两半输出, 这里用真实字块边界。
gb = cv2.cvtColor(big, cv2.COLOR_RGB2GRAY)
col_dark = (gb < 160).sum(axis=0)
blocks, in_b, b0 = [], False, 0
for x, d in enumerate(col_dark):
if d > 2 and not in_b:
in_b, b0 = True, x
elif d <= 2 and in_b:
in_b = False
if x - b0 >= 2:
blocks.append([b0, x])
if in_b:
blocks.append([b0, len(col_dark)])
if not blocks:
continue
char_h = big.shape[0] # 旋转后字高 ≈ 原列宽, 作为切字尺子
gap_min = max(2, int(char_h * 0.2)) # 断笔间隙 <= 该值 -> 同字, 合并
merged = [blocks[0]]
for blk in blocks[1:]:
if blk[0] - merged[-1][1] <= gap_min:
merged[-1][1] = blk[1]
else:
merged.append(blk)
# 字宽尺子: 合并后块宽的中位数(断笔已合并, 单字块为主)。
# 竖排字宽不等于列宽(Det 框可能框住1~2列), 不能用 char_h 直接当字宽;
# 连笔列会合并成单个巨型块污染中位数, 字宽也不可能超过列宽, 故取上界
widths = sorted(b[1] - b[0] for b in merged)
w_med = widths[len(widths) // 2] if widths else char_h
w_med = max(2, min(w_med, int(0.9 * char_h)))
# 断笔残块并入相邻块: 飞白/横笔间隙切出的残块明显窄于字宽,
# 合并后按字宽整块识别, 避免半个字独立成框
i = 0
while i < len(merged):
if merged[i][1] - merged[i][0] < 0.65 * w_med:
left_gap = merged[i][0] - merged[i - 1][1] if i > 0 else float("inf")
right_gap = merged[i + 1][0] - merged[i][1] if i < len(merged) - 1 else float("inf")
if left_gap <= right_gap and i > 0:
merged[i - 1][1] = merged[i][1]
del merged[i]
i -= 1
elif i < len(merged) - 1:
merged[i][1] = merged[i + 1][1]
del merged[i + 1]
else:
i += 1
else:
i += 1
# 超宽块拆开: 合并后宽 > 1.4×字宽说明粘了多个字, 在最大间隙处递归拆
def split_wide(blk, subs):
if blk[1] - blk[0] <= 1.4 * w_med or len(subs) < 2:
return [blk]
gaps = [(subs[i + 1][0] - subs[i][1], i) for i in range(len(subs) - 1)]
gi = max(range(len(gaps)), key=lambda i: gaps[i][0])
if gaps[gi][0] < 0.2 * w_med:
return [blk]
cut = (subs[gi][1] + subs[gi + 1][0]) // 2
left = [blk[0], cut]; right = [cut, blk[1]]
return (split_wide(left, [s for s in subs if s[1] <= cut])
+ split_wide(right, [s for s in subs if s[0] >= cut]))
final_blocks = []
for blk in merged:
subs = [s for s in blocks if s[0] >= blk[0] - 1 and s[1] <= blk[1] + 1]
final_blocks.extend(split_wide(blk, subs))
new_items = []
for bx0, bx1 in final_blocks:
if bx1 - bx0 < 8:
continue
# 连笔/倾斜导致投影切不开的超宽块: 按字宽等份切开, 逐份识别
if bx1 - bx0 > 1.4 * w_med:
n = max(2, round((bx1 - bx0) / w_med))
cuts = [bx0 + (bx1 - bx0) * i // n for i in range(n + 1)]
blocks_2 = [[cuts[i], cuts[i + 1]] for i in range(n)]
else:
blocks_2 = [[bx0, bx1]]
for cx0, cx1 in blocks_2:
r = self.engine(big[:, cx0:cx1], use_det=False, use_cls=False, use_rec=True)
txt = "".join(r.txts) if r.txts else ""
if not txt:
continue
# 逆时针旋转90°后: big 的 x' 方向(宽2×chh)对应 crop 的 y 方向,
# y' 方向(高2×chw)对应 crop 的 x 方向。
# crop(x_c, y_c) = (chw-1-y'/2, x'/2), 字块占满 y' 全高 -> x_c 全宽
oy1 = y1 + cx0 / 2; oy2 = y1 + cx1 / 2
ox1 = x1; ox2 = x1 + chw - 1
coords = [[int(ox1), int(oy1)], [int(ox2), int(oy1)],
[int(ox2), int(oy2)], [int(ox1), int(oy2)]]
cxx, cyy = center(coords)
if any(abs(cxx - center(c)[0]) <= 10 and abs(cyy - center(c)[1]) <= 10
for _, _, c in orig_items + new_items):
continue
new_items.append((txt, 0.5, coords))
if new_items:
orig_items = orig_items + new_items
logger.info(f"竖排长框旋转补充 {len(new_items)} 字 (列x~{cx:.0f})")
# 主流程整列条目若被逐字补充充分覆盖(>=2个字且覆盖>=50%面积),
# 则丢弃, 只保留字级框, 避免前端画出一堆叠加框
n_new = len(new_items)
if n_new:
main_part = orig_items[:len(orig_items) - n_new]
new_coords = [c for _, _, c in new_items]
def covered(item):
t, s, c = item
if not isinstance(t, str) or not t:
return False
x0, y0 = min(p[0] for p in c), min(p[1] for p in c)
x1, y1 = max(p[0] for p in c), max(p[1] for p in c)
area = (x1 - x0) * (y1 - y0)
if area <= 0:
return False
hits, hit_area = 0, 0
for nc in new_coords:
nx0, ny0 = min(p[0] for p in nc), min(p[1] for p in nc)
nx1, ny1 = max(p[0] for p in nc), max(p[1] for p in nc)
ix0, iy0 = max(x0, nx0), max(y0, ny0)
ix1, iy1 = min(x1, nx1), min(y1, ny1)
if ix1 > ix0 and iy1 > iy0:
hits += 1
hit_area += (ix1 - ix0) * (iy1 - iy0)
return hits >= 2 and hit_area >= 0.5 * area
main_part = [item for item in main_part if not covered(item)]
# 列内相对过滤: 明显矮于同列中位的 0.5 分字块是残笔/等份碎片, 删除
if len(new_items) >= 3:
by_col = {}
for it in new_items:
cxx = sum(p[0] for p in it[2]) / len(it[2])
by_col.setdefault(round(cxx / 30), []).append(it)
for key, its in by_col.items():
if len(its) < 3:
continue
hs = sorted(max(p[1] for p in c) - min(p[1] for p in c) for _, _, c in its)
med = hs[len(hs) // 2]
by_col[key] = [it for it in its
if max(p[1] for p in it[2]) - min(p[1] for p in it[2]) >= med * 0.5]
new_items = [it for its in by_col.values() for it in its]
orig_items = main_part + new_items
result.word_results = tuple(
(t, s, c) if isinstance(t, str) else (t[0], t[1], t[2])
for t, s, c in orig_items
)
return result
def process_with_vertical(self, image: np.ndarray) -> RapidOCROutput:
"""标准 OCR + 边缘/竖排文字补充通道(纯 CV + 长框旋转识别,合并检测框与文本)"""
result = self.process(image)
if not self.enable_vertical:
return result
main_boxes = list(result.boxes) if result.boxes is not None else []
extra = self._cv_supplement(image, main_boxes)
if extra:
merged = []
for box, text in extra:
cx, cy = box[:, 0].mean(), box[:, 1].mean()
# 中心落在主流程框内则跳过
if any(b[:, 0].min() - 4 <= cx <= b[:, 0].max() + 4
and b[:, 1].min() - 4 <= cy <= b[:, 1].max() + 4 for b in main_boxes):
continue
# 补充块之间去重(多阈值可能检出同一区域)
if any(abs(cx - mx) < 15 and abs(cy - my) < 15 for _, _, _, mx, my in merged):
continue
coords = [[int(p[0]), int(p[1])] for p in box]
merged.append((text, 1.0, coords, cx, cy))
if merged:
existing = list(result.word_results) if result.word_results else []
result.word_results = tuple(existing + [(t, s, c) for t, s, c, _, _ in merged])
logger.info(f"竖排通道补充 {len(merged)} 条: {[t for t, _, _, _, _ in merged]}")
# 竖排长框旋转补充(恢复超长行识别丢字)
result = self._vertical_longbox_supplement(image, result)
return result