Files

527 lines
24 KiB
Python
Raw Permalink Normal View History

2026-08-31 16:41:12 +08:00
"""
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