Spaces:
Sleeping
Sleeping
| """ | |
| 纯文字 OCR 识别系统 (基于 V4.24 简化) | |
| 移除: 公式检测 (MFD) / 公式识别 (MFR) | |
| 保留: 文本检测 (DBNet) / 文本识别 (CRNN) / 文本行合并 | |
| """ | |
| import cv2 | |
| import numpy as np | |
| import onnxruntime as ort | |
| import pyclipper | |
| from shapely.geometry import Polygon | |
| import os | |
| import math | |
| # ========================================== | |
| # 1. 基础工具与后处理 (保持 V4.24 修复版) | |
| # ========================================== | |
| class DBPostProcess: | |
| """DBNet 后处理""" | |
| def __init__(self, thresh=0.3, box_thresh=0.6, max_candidates=1000, unclip_ratio=1.5): | |
| self.thresh = thresh | |
| self.box_thresh = box_thresh | |
| self.max_candidates = max_candidates | |
| self.unclip_ratio = unclip_ratio | |
| self.min_size = 3 | |
| def __call__(self, pred, shape_list): | |
| pred = pred[0, 0, :, :] | |
| segmentation = pred > self.thresh | |
| boxes_batch = [] | |
| scores_batch = [] | |
| mask = (segmentation * 255).astype(np.uint8) | |
| contours, _ = cv2.findContours(mask, cv2.RETR_LIST, cv2.CHAIN_APPROX_SIMPLE) | |
| src_h, src_w, ratio_h, ratio_w = shape_list | |
| for contour in contours: | |
| if len(contour) < 2: continue | |
| # 第一次调用 get_mini_boxes | |
| box, sside = self.get_mini_boxes(contour) | |
| if sside < self.min_size: continue | |
| points = np.array(box) | |
| score = self.box_score_fast(pred, points) | |
| if score < self.box_thresh: continue | |
| # unclip 返回的是列表 (list of points) | |
| box = self.unclip(points) | |
| if box is None: continue | |
| # 第二次调用 get_mini_boxes | |
| box, sside = self.get_mini_boxes(box) | |
| if sside < self.min_size + 2: continue | |
| box = np.array(box) | |
| box[:, 0] = np.clip(np.round(box[:, 0] / ratio_w), 0, src_w) | |
| box[:, 1] = np.clip(np.round(box[:, 1] / ratio_h), 0, src_h) | |
| boxes_batch.append(box.astype(np.int32)) | |
| scores_batch.append(score) | |
| return boxes_batch, scores_batch | |
| def unclip(self, box): | |
| poly = Polygon(box) | |
| if poly.length == 0: return None | |
| distance = poly.area * self.unclip_ratio / poly.length | |
| offset = pyclipper.PyclipperOffset() | |
| offset.AddPath(box, pyclipper.JT_ROUND, pyclipper.ET_CLOSEDPOLYGON) | |
| expanded = offset.Execute(distance) | |
| return expanded[0] if len(expanded) > 0 else None | |
| def get_mini_boxes(self, contour): | |
| # 强制转换为 numpy array (修复 OpenCV 报错) | |
| if not isinstance(contour, np.ndarray): | |
| contour = np.array(contour, dtype=np.float32) | |
| bounding_box = cv2.minAreaRect(contour) | |
| points = sorted(list(cv2.boxPoints(bounding_box)), key=lambda x: x[0]) | |
| index_1, index_2, index_3, index_4 = 0, 1, 2, 3 | |
| if points[1][1] > points[0][1]: index_1, index_4 = 0, 1 | |
| else: index_1, index_4 = 1, 0 | |
| if points[3][1] > points[2][1]: index_2, index_3 = 2, 3 | |
| else: index_2, index_3 = 3, 2 | |
| box = [points[index_1], points[index_2], points[index_3], points[index_4]] | |
| return box, min(bounding_box[1]) | |
| def box_score_fast(self, bitmap, _box): | |
| h, w = bitmap.shape[:2] | |
| box = _box.copy() | |
| xmin = np.clip(np.floor(box[:, 0].min()).astype(int), 0, w - 1) | |
| xmax = np.clip(np.ceil(box[:, 0].max()).astype(int), 0, w - 1) | |
| ymin = np.clip(np.floor(box[:, 1].min()).astype(int), 0, h - 1) | |
| ymax = np.clip(np.ceil(box[:, 1].max()).astype(int), 0, h - 1) | |
| if xmax <= xmin or ymax <= ymin: return 0 | |
| mask = np.zeros((ymax - ymin + 1, xmax - xmin + 1), dtype=np.uint8) | |
| box[:, 0] = box[:, 0] - xmin | |
| box[:, 1] = box[:, 1] - ymin | |
| cv2.fillPoly(mask, box.reshape(1, -1, 2).astype(np.int32), 1) | |
| return cv2.mean(bitmap[ymin:ymax + 1, xmin:xmax + 1], mask)[0] | |
| def get_rotate_crop_image(img, points): | |
| points = np.array(points, dtype=np.float32) | |
| width = int(max(np.linalg.norm(points[0] - points[1]), np.linalg.norm(points[2] - points[3]))) | |
| height = int(max(np.linalg.norm(points[0] - points[3]), np.linalg.norm(points[1] - points[2]))) | |
| pts_std = np.float32([[0, 0], [width, 0], [width, height], [0, height]]) | |
| M = cv2.getPerspectiveTransform(points, pts_std) | |
| dst_img = cv2.warpPerspective(img, M, (width, height), borderMode=cv2.BORDER_REPLICATE, flags=cv2.INTER_CUBIC) | |
| if dst_img.shape[0] * 1.0 / dst_img.shape[1] >= 1.5: | |
| dst_img = np.rot90(dst_img) | |
| return dst_img | |
| # ========================================== | |
| # 2. 文本检测器 (DBNet) | |
| # ========================================== | |
| class TextDetector: | |
| def __init__(self, model_path, use_gpu=True): | |
| print(f"Loading DBNet model from: {model_path}") | |
| providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if use_gpu else ['CPUExecutionProvider'] | |
| self.session = ort.InferenceSession(model_path, providers=providers) | |
| self.input_name = self.session.get_inputs()[0].name | |
| self.post_process = DBPostProcess(thresh=0.3, box_thresh=0.5, unclip_ratio=1.5) | |
| def predict(self, img, limit_side_len=1280): | |
| h, w = img.shape[:2] | |
| ratio = 1.0 | |
| if max(h, w) > limit_side_len: | |
| ratio = float(limit_side_len) / max(h, w) | |
| resize_h = int(h * ratio) | |
| resize_w = int(w * ratio) | |
| resize_h = max(int(round(resize_h / 32) * 32), 32) | |
| resize_w = max(int(round(resize_w / 32) * 32), 32) | |
| ratio_h = resize_h / float(h) | |
| ratio_w = resize_w / float(w) | |
| img_resized = cv2.resize(img, (resize_w, resize_h)) | |
| # 显式 float32 运算 | |
| img_norm = img_resized.astype(np.float32) / 255.0 | |
| img_norm -= np.array([0.485, 0.456, 0.406], dtype=np.float32) | |
| img_norm /= np.array([0.229, 0.224, 0.225], dtype=np.float32) | |
| img_input = img_norm.transpose((2, 0, 1))[np.newaxis, :] | |
| img_input = img_input.astype(np.float32) | |
| outputs = self.session.run(None, {self.input_name: img_input}) | |
| boxes, scores = self.post_process(outputs[0], [h, w, ratio_h, ratio_w]) | |
| return boxes, scores | |
| # ========================================== | |
| # 3. 文本识别器 (CRNN) | |
| # ========================================== | |
| class TextRecognizer: | |
| def __init__(self, model_path, label_path, use_gpu=True): | |
| print(f"Loading CRNN model from: {model_path}") | |
| providers = ['CUDAExecutionProvider', 'CPUExecutionProvider'] if use_gpu else ['CPUExecutionProvider'] | |
| self.session = ort.InferenceSession(model_path, providers=providers) | |
| self.input_name = self.session.get_inputs()[0].name | |
| self.character = self.load_dict(label_path) | |
| self.rec_image_height = 48 | |
| def load_dict(self, label_path): | |
| with open(label_path, 'r', encoding='utf-8') as f: | |
| character = [line.strip('\n') for line in f.readlines()] | |
| return ['blank'] + character + [' '] | |
| def predict(self, img_list): | |
| if len(img_list) == 0: return [] | |
| imgH = self.rec_image_height | |
| max_wh_ratio = max([w / h for h, w in [img.shape[:2] for img in img_list]] + [320.0 / imgH]) | |
| batch_input = [] | |
| for img in img_list: | |
| h, w = img.shape[:2] | |
| imgW = int(imgH * max_wh_ratio) | |
| ratio = w / float(h) | |
| resized_w = int(math.ceil(imgH * ratio)) if math.ceil(imgH * ratio) <= imgW else imgW | |
| resized_image = cv2.resize(img, (resized_w, imgH)).astype(np.float32) | |
| resized_image = (resized_image.transpose((2, 0, 1)) / 255.0 - 0.5) / 0.5 | |
| padding_im = np.zeros((3, imgH, imgW), dtype=np.float32) | |
| padding_im[:, :, 0:resized_w] = resized_image | |
| batch_input.append(padding_im[np.newaxis, :]) | |
| input_tensor = np.concatenate(batch_input, axis=0) | |
| outputs = self.session.run(None, {self.input_name: input_tensor})[0] | |
| preds_idx = outputs.argmax(axis=2) | |
| results = [] | |
| for idxs in preds_idx: | |
| res = "" | |
| last = 0 | |
| for idx in idxs: | |
| if idx != last and idx != 0 and idx < len(self.character): | |
| res += self.character[idx] | |
| last = idx | |
| results.append(res) | |
| return results | |
| # ========================================== | |
| # 4. 文本行合并器 (TextMerger) - 优化版 | |
| # ========================================== | |
| class TextMerger: | |
| def __init__(self, y_thresh=20): | |
| self.y_thresh = y_thresh | |
| def merge(self, text_results): | |
| """ | |
| 优化合并逻辑:确保严格的'从左到右,从上到下'顺序 | |
| 使用基于重叠度的行聚类,而非简单的线性阈值 | |
| """ | |
| if not text_results: return [] | |
| # 1. 预计算几何信息 | |
| for item in text_results: | |
| box = item['box'] | |
| y_min, y_max = np.min(box[:, 1]), np.max(box[:, 1]) | |
| x_min = np.min(box[:, 0]) | |
| item['y_min'] = y_min | |
| item['y_max'] = y_max | |
| item['y_center'] = (y_min + y_max) / 2 | |
| item['x_min'] = x_min | |
| item['height'] = y_max - y_min | |
| # 2. 初始按 y_min 排序,作为基准选择顺序 | |
| text_results.sort(key=lambda x: x['y_min']) | |
| lines = [] | |
| # 3. 迭代聚类 | |
| while len(text_results) > 0: | |
| # 取出当前最顶部的框作为"行基准" | |
| base = text_results.pop(0) | |
| current_line = [base] | |
| remains = [] | |
| base_y_center = base['y_center'] | |
| base_height = base['height'] | |
| for item in text_results: | |
| # 判定同行:中心点 Y 轴距离小于一定阈值 | |
| avg_height = (base_height + item['height']) / 2 | |
| dynamic_thresh = max(self.y_thresh, avg_height * 0.5) | |
| y_diff = abs(item['y_center'] - base_y_center) | |
| if y_diff < dynamic_thresh: | |
| current_line.append(item) | |
| else: | |
| remains.append(item) | |
| text_results = remains | |
| # 4. 行内排序:严格按 X 轴 (从左到右) | |
| current_line.sort(key=lambda x: x['x_min']) | |
| lines.append(current_line) | |
| # 5. 行间排序:虽已大致有序,可再按每行第一个元素的 Y 重新微调 | |
| lines.sort(key=lambda line: line[0]['y_min']) | |
| # 6. 拼接输出 | |
| final_output = [] | |
| for line in lines: | |
| # 简单空格连接 | |
| line_text = " ".join([x['content'] for x in line]) | |
| # 清理多余空格 | |
| line_text = " ".join(line_text.split()) | |
| final_output.append(line_text) | |
| return final_output | |
| # ========================================== | |
| # 5. 主处理器 (PureOCRSystem) | |
| # ========================================== | |
| class PureOCRSystem: | |
| def __init__(self, det_path, rec_path, key_path, use_gpu=True): | |
| print("\n" + "="*50) | |
| print("初始化纯文字识别系统 (无公式)") | |
| print("="*50) | |
| self.detector = TextDetector(det_path, use_gpu=use_gpu) | |
| self.recognizer = TextRecognizer(rec_path, key_path, use_gpu=use_gpu) | |
| self.merger = TextMerger(y_thresh=20) | |
| def process(self, img_path): | |
| """处理本地图片文件""" | |
| img = cv2.imdecode(np.fromfile(img_path, dtype=np.uint8), cv2.IMREAD_COLOR) | |
| if img is None: | |
| print(f"❌ 无法读取图像: {img_path}") | |
| return None | |
| return self._process_image(img) | |
| def process_image(self, img): | |
| """处理 numpy 数组格式的图片(FastAPI 使用)""" | |
| return self._process_image(img) | |
| def _process_image(self, img): | |
| """内部处理函数""" | |
| # 1. 文本检测 (直接在原图上进行) | |
| print("Step 1: 检测文本...") | |
| text_boxes, scores = self.detector.predict(img) | |
| print(f" 找到 {len(text_boxes)} 个文本框") | |
| # 2. 文本识别 | |
| print("Step 2: 识别文本...") | |
| text_results = [] | |
| if len(text_boxes) > 0: | |
| # 简单按Y坐标初排,方便批量处理 | |
| sorted_indices = np.argsort(np.array(text_boxes)[:, 0, 1]) | |
| text_boxes = [text_boxes[i] for i in sorted_indices] | |
| crop_list = [] | |
| for box in text_boxes: | |
| crop = get_rotate_crop_image(img, box) | |
| crop_list.append(crop) | |
| raw_texts = self.recognizer.predict(crop_list) | |
| for i, box in enumerate(text_boxes): | |
| text_results.append({ | |
| 'box': box.astype(float), | |
| 'content': raw_texts[i], | |
| 'type': 'text' | |
| }) | |
| # 3. 行合并 | |
| print("Step 3: 结果合并 (从左到右,从上到下)...") | |
| final_lines = self.merger.merge(text_results) | |
| return final_lines |