527 lines
24 KiB
Python
527 lines
24 KiB
Python
"""
|
||
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
|
||
|