""" 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