ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Pytorch OCR工具库实战:从DBNet+CRNN到ONNX部署

Pytorch OCR工具库实战:从DBNet+CRNN到ONNX部署 简介这是一份基于Pytorch框架的OCR工具库资源包面向计算机视觉初学者、算法研究人员以及课程设计开发者。它集成了常见的文字检测与识别算法支持从图像预处理到模型推理的完整流程可帮助用户快速搭建OCR实验环境并进行二次开发。压缩包共687个文件总大小约9.41MB其中包含大量Python脚本约415个py用于核心算法实现与训练推理YAML配置用于模型参数管理TXT文档用于样本列表与说明JPG/PNG图片作为测试样例另有少量Pytorch权重文件pkl和C/CUDA扩展源码适合在此基础上开展目标检测、文本识别等方向的项目实践。该资源在CSDN已有238人浏览学习适合作为课程设计参考或算法入门辅助资料。压缩包内目录组织清晰如roi_align_rotated等扩展模块以及测试图片、手册文档便于用户对照源码理解实现细节。1. 基于Pytorch的OCR工具库文字检测和识别算法怎么组织才不翻车批量做票据识别、卡证识别或扫描件归档时很多团队第一步会去接现成OCR引擎或云端API等遇到私有化部署、中文长文本、表格混排时才发现可控比开箱即用重要。基于Pytorch自建一个OCR工具库听上去只是把文字检测和识别算法串起来真正写代码时会发现难点不在模型本身而在数据标注格式、真值生成、后处理和训练参数之间的匹配。这篇文章顺着一条最稳的落地路径讲选什么检测和识别算法怎么用Pytorch把它们拼成可复用库踩过哪些常见的坑以及最后怎么把模型导出成ONNX交给业务侧。2. 文字检测和识别算法的选型DBNet加CRNN为什么是默认起点2.1 两阶段流水线的价值检测和识别分开能换来的可控性OCR工具库最常用的组织方式是分两段文字检测在整图上定位文本区域输出矩形框或多边形文字识别把裁剪区域转成字符串。为什么两阶段还是主流而不是像Mask TextSpotter那样端到端一次输出因为训练目标和数据依赖不同。检测模型只需要知道哪里是字不需要理解语言顺序识别模型只需要看懂这一行字是什么不需要处理版面。拆开之后哪一环弱就单独补数据、单独调参数产品或运维也能清晰描述问题出在检测还是识别。端到端模型的精度上限未必低但调参和排错成本高工程上除非有特殊强需求一般不会优先选。两阶段配合得当在中文票据、物流面单、银行流水这些常见场景里已经完全够用。Pipeline层面还有一个好处识别网络可以反复复用。检测模型从整图定位出文本框后识别模型见到的始终是固定高度的文本行。这样训练识别模型的数据可以单独构造甚至可以先用公开数据预训练再用业务数据微调。2.2 检测算法选型DBNet、EAST和PSENet分别解决什么问题在基于Pytorch的OCR工具库里检测部分最常见的三个算法是DBNet、EAST和PSENet。DBNetDifferentiable Binarization的特点是把传统后处理里的二值化阈值变成一个可学习的 threshold map网络同时预测 probability map 和 threshold map二者组合出近似二值图。它解决了传统分割后处理断开连接的问题对长文本和密排文本的还原能力出色推理时只需要一个简单后处理因此速度也快。在多数票据场景中我一般直接用DBNet作为底座训练稳定不用加额外的角度分支。EAST直接回归到旋转矩形框刚好符合一行文字的形状先验但EAST对弯曲文本和竖排文本支持弱且需要自己维护旋转框的编码方式数据预处理比DBNet复杂。PSENet用渐进式尺度扩展解决粘连文本的分离在弯曲文本、任意四边形甚至弓形文本上有优势但多个尺度的后处理循环会拖慢推理速度。算法输出形式长文本弯曲/竖排推理成本典型场景DBNet多边形/矩形好一般低票据、文档、屏幕EAST旋转矩形好弱低横排密集文本PSENet多边形中偏上好中高不规则版面2.3 识别算法选型CRNN、SVTR与Transformer解码器的取舍识别网络通常接收高度32像素左右的文本行图。CRNN的网络结构是CNN提取视觉特征、RNN序列建模、CTC对齐它对整行标注非常友好不需要逐字符标注因而是开源OCR工具库里出现频率最高的方案。中文词表动辄几千字CTC在这种大词表上依然稳定这是它到今天还不被淘汰的原因。SVTR是纯视觉Transformer结构用全局注意力替代RNN的时序建模在英文场景的准确率和速度都优于CRNN也能搭配CTC。中英文混合场景下SVTR加轻量语言模型可以做得更准但显存占用高一些部署时还需照顾Transformer的编译优化。还有一类基于Transformer解码器的方法比如SATRN、ABINet对自然场景里的不规则文本效果好却需要一个可学习的注意力对齐过程遇到长文本行容易发生解码漂移而且训练数据最好有字符级标注。工具库要服务一个长期演进的团队我更推荐默认走DBNet CRNN/SVTR CTC的组合先把大部分业务跑通再按Badcase决定要不要引入语言模型。实际上选型不必一步到位。Pytorch工具库的价值就在于检测head、backbone、识别head都是可替换组件换算法只影响对应模块不影响数据管道和后处理接口。3. 搭建基于Pytorch的OCR工具库数据管道、模型组装与训练接口3.1 最小目录结构与核心抽象先说工具库的形态。它不该是一堆训练脚本而应该像一辆组装好的车外部调用方只面对几个类的方法内部各模块可替换。ocr_lib/ configs/ det_db.yaml rec_crnn.yaml core/ dataset.py # 数据读取、增广、标签生成 transforms.py # 归一化、缩放、shrink mask det/ models.py # DBNet/EAST 检测网络 loss.py # DB loss / L1 loss postprocess.py # 多边形的还原 rec/ models.py # CRNN/SVTR 识别网络 loss.py # CTC loss / CE loss postprocess.py # 从输出序列解码出文本 pipeline.py # 对外提供的 OCRPipeline对外给业务方使用的接口类似这样from ocr_lib.pipeline import OCRPipeline det_args {model_path: checkpoints/det_db.pt, conf_thresh: 0.5} rec_args {model_path: checkpoints/rec_crnn.pt, max_len: 25} ocr OCRPipeline(det_args, rec_args) results ocr.predict(image_pathinvoice.jpg) # results: [{box: [[x1,y1],[x2,y2],[x3,y3],[x4,y4]], # text: 购货方, conf: 0.96}, ...]这里predict内部先做检测、再做识别。工具库的一个关键设计是检测和识别各自持有独立的预处理和后处理函数OCR Pipeline 里通过文本行裁剪来衔接两段任何一段升级都不影响另一段的调用方式。3.2 数据组织从标注文件到可训练样本检测训练一般用ICDAR格式每张图对应一个同名的txt或json按行记录多边形的四个角点格式形如x1,y1,x2,y2,x3,y3,x4,y4,文本内容。要注意训练DBNet需要的不只是原始四边形还要按标注多边形生成三个与输入同尺寸的2D标签矩阵分别是概率图、阈值图和shrink mask。# src/ocr_lib/core/generate_label.py import numpy as np from shapely.geometry import Polygon, box def make_db_label(poly, img_h, img_w, shrink_ratio0.4): # poly: 四点多边形 [[x1,y1],[x2,y2],[x3,y3],[x4,y4]] polygon Polygon(poly) distance polygon.area * (1 - shrink_ratio) / polygon.length shrink_poly polygon.buffer(-distance) # 初始化概率图和阈值图 pgt np.zeros((img_h, img_w), dtypenp.uint8) tgt np.zeros((img_h, img_w), dtypenp.float32) # 将缩小后多边形内的像素填充为1构成概率图真值 pgt fill_polygon(shrink_poly, pgt, 1.0) tgt fill_polygon(polygon, tgt, 1.0) return pgt, tgt def fill_polygon(polygon, canvas, value): import cv2 coords np.array(polygon.exterior.coords, dtypenp.int32) cv2.fillPoly(canvas, [coords], value) return canvasshrink_ratio决定检测真值的收缩程度过大会让检测框被吃进去过小会让相邻文本行在loss里互相挤压。DBNet的改进点在于使用可学习的threshold图让模型自己判断哪里应该二值化从而放宽了对固定收缩值的要求但标签生成里的shrink_ratio仍然是先验约束。识别训练数据比检测简单每行是一张固定高度的灰度图高度32或48配一个字符串。训练时要把字符串编码成vocab的索引序列并在每个样本前面拼上空白token以便CTC对齐。如果业务文本行超长导入时先统计长度分布把最长的那一批直接截断避免为了几行超长数据放大整批的max_len那会显著增加CTC序列中blank的占比、拉慢收敛。3.3 用Pytorch组装检测和识别模型检测模型最简洁的组合是 ResNet18/34 作为骨干FPN汇聚多尺度特征再接DB head。一个能跑通训练的示例import torch import torch.nn as nn class DBHead(nn.Module): def __init__(self, in_channels, hidden256): super().__init__() self.conv1 nn.Conv2d(in_channels, hidden, 3, padding1) self.bn1 nn.BatchNorm2d(hidden) self.prob nn.Conv2d(hidden, 1, 1) self.thresh nn.Conv2d(hidden, 1, 1) def forward(self, features): x torch.relu(self.bn1(self.conv1(features))) return self.prob(x), self.thresh(x) def detect_loss(pred_prob, pred_thresh, pgt_prob, pgt_thresh): bce nn.BCEWithLogitsLoss() prob_loss bce(pred_prob, pgt_prob.float()) mask pgt_prob 1.0 thresh_loss nn.L1Loss()(pred_thresh[mask], pgt_thresh[mask]) return prob_loss 10.0 * thresh_loss说明DBHead同时输出概率图和阈值图两者在前向里合成近似二值图。训练时概率图用BCE阈值图用L1第二个损失的系数一般给到10.0阈值分支的准确性直接影响推理阶段二值化的稳定性。损失只在文本区域内计算阈值损失是因为背景区域的阈值没有明确真值强行约束反而会让网络震荡。识别端给一个CRNN的极简实现框架class CRNN(nn.Module): def __init__(self, backbone_out512, num_class6623, hidden256): super().__init__() self.backbone nn.Sequential( nn.Conv2d(1, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.lstm nn.LSTM(backbone_out, hidden, bidirectionalTrue, batch_firstTrue) self.fc nn.Linear(hidden * 2, num_class) self.ctc_loss nn.CTCLoss(blank0, zero_infinityTrue) def forward(self, x, labelsNone, label_lensNone): x self.backbone(x) # B, C, H, W b, c, h, w x.shape x x.view(b, c * h, w).permute(0, 2, 1) # B, W, C x, _ self.lstm(x) out self.fc(x) # B, W, num_class if labels is not None: out_lens torch.full((b,), w, dtypetorch.long) loss self.ctc_loss(out.permute(1, 0, 2), labels, out_lens, label_lens) return out, loss return out注意这里第0类留给CTC的blank损失函数里blank0必须和vocab编码一致。zero_infinityTrue防止某个batch里出现空标注导致loss变成负无穷。CNN输出后把H和C合并成特征维度W作为序列长度配合双向LSTM建模横向的字符顺序。这是CRNN压缩到最核心的数据流实际工程中backbone会更深、带BatchNorm和dropout但序列特征的变换过程是同一套。4. 训练与调参Pytorch环境搭建和OCR专用参数怎么设4.1 Anaconda配置Pytorch环境CPU/GPU版本和下载慢怎么办OCR工具库对训练机的要求是Pytorch能正常跑GPU对推理机则往往只要CPU或轻量GPU。跨机器复现时最省心的办法是Anaconda建一个独立环境conda create -n ocr python3.10 -y conda activate ocr pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 pip install opencv-python shapely pyyaml tqdm如果速度不理想把pip的index换到清华或阿里云的镜像站例如pip install torch torchvision -i https://pypi.tuna.tsinghua.edu.cn/simpletorch的wheel体积大耐心等。装完后用下面这行验证GPU是否真的可用python -c import torch; print(torch.__version__, torch.cuda.is_available(), torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU)torch.cuda.is_available()返回True只代表编译好的CUDA版本能加载真正训练前还要确认torch.device(cuda)上的张量能跑一次乘法。相比之下更值得检查的是Pytorch的编译版本和显卡驱动是否匹配老驱动不更新时会看到 CUDA error: no kernel image is available for execution on the device这时需要退回对应的pytorch wheel。此外同一台机器上如果先前装过TensorFlow它的CUDA runtime和Pytorch冲突建议把两个环境彻底隔离。4.2 OCR训练参数表batch size、学习率和shrink ratio同样一组基于Pytorch的OCR训练代码不同数据规模下参数差异很大。以下表格是固定图像尺寸、单机训练时的默认值单卡或CPU则按比例缩小batch并调低学习率参数检测DBNet识别CRNN调整方向输入尺寸960x640320x32识别高度固定宽度按最长文本动态设置batch_size81664128显存不足时先降batch别动学习率学习率1e-43e-4检测收敛慢OCR用warmup 500步shrink_ratio0.4~0.6-相邻文本近时调小文本稀疏时调大box_thresh0.3~0.5-推理时用于过滤低置信框unclip_ratio1.5~2.0-控制检测框向外扩多少OCR阶段需要留边max_len-25~50中文票据常用25超长文本调到40以上一个经常被忽略的匹配问题是输入尺度。检测阶段标注多边形来自原始图坐标训练时resize到960x640推理时新图的resize方式必须与训练一致通常使用等比例缩放并对短边pad而不是直接拉伸。识别阶段由于高度固定32且宽度不固定Pytorch的nn.CTCLoss要求输入序列长度与标签长度在概率上有意义宽度被两次池化后序列长度是原始宽度的四分之一训练时就不能随意的把图resize到任意宽度否则字符会被压扁或拉伸。CTC loss训练有个常见误用是padding导致的对齐混乱。CTCLoss给每个样本独立的标签长度本批里不同长度的标注用label_lens区分千万不要给短文本统一补空格那会让CTC统计出很多错误的loss项。如果看到一个OCR工具库训练不收敛、准确率在个位数徘徊优先怀疑这一处。4.3 评估检测F1与识别整串准确率的计算方法训练中定期评估比只看train loss可靠得多。检测评估是对预测框和真值框做IoU匹配IoU高于0.5记为命中统计Precision、Recall与F1。识别的任务则更接近分类准确率但工业里真正有用的是整串完全匹配的准确率Strict Accuracy以及带编辑距离的归一化准确率。def evaluate_recognition(model, val_loader, decode_func, devicecuda): correct 0 total 0 normalized_edit 0.0 for images, labels, label_lens in val_loader: images images.to(device) out model(images) # B, T, num_class preds decode_func(out) # 返回每个样本的文本 for pred, target in zip(preds, labels): target_len int(label_lens[total]) target_text decode_target(target, target_len) total 1 edit edit_distance(pred, target_text) normalized_edit 1 - edit / max(len(target_text), 1) if pred target_text: correct 1 return correct / total, normalized_edit / totaldecode函数要把CTC输出里的重复字符和blank去掉这一层解码逻辑属于工具库的公共模块。对测试集建议用英文公开数据集IIIT5K或SVT做一次通用性测试再用自研业务数据做回归测试避免在单一数据集上过拟合到某种字体或版式。编辑距离用python-Levenshtein包计算比手写循环快一个数量级。5. 推理与部署把Pytorch模型导出成ONNX的验证技巧5.1 检测和识别串成一条可复用的推理管线训练完成后常见流程是先用可视化脚本看检测框是否正确框住文本再单独抽识别badcase。可视化脚本不能省它往往比一切日志更快发现数据标注错位例如多边形的点序错误会导致画出来的框交叉这类问题在训练初期很常见。python tools/visualize_ocr.py \ --det checkpoint/det_db.pt \ --rec checkpoint/rec_crnn.pt \ --img data/sample_invoice.jpg \ --save_dir viz_out脚本的输出是一张标注好框和识别文本的图片配合命令行打印的[box, text, conf]表格能直观定位检测框重了但识别文本对识别文本漏字但框准的两类典型故障。建议每个迭代轮次都随机抽50张badcase照片归档后续标注这批数据做主动学习收益远大于不停调学习率。5.2 用ONNX导出把推理服务与Pytorch解耦工具库最终往往要部署成HTTP服务或嵌入桌面程序。Pytorch官方torch.onnx.export可以把检测和识别模型分别导出导出后前后处理仍留在Python层这样既能享受ONNX Runtime的速度优势又不至于把复杂的图像预处理和DB后处理锁死在导出图里。import torch from ocr_lib.det.models import build_db_model model build_db_model(backboneresnet18, pretrainedFalse) model.load_state_dict(torch.load(checkpoint/det_db.pt, map_locationcpu)) model.eval() dummy torch.randn(1, 3, 640, 960) torch.onnx.export( model, dummy, det_db.onnx, input_names[input], output_names[prob, thresh], dynamic_axes{input: {0: batch, 2: height, 3: width}}, opset_version12 )动态轴dynamic_axes是必填而非可选项。OCR服务里请求图片尺寸差异极大若导出时把图片长宽固定为960x640推理阶段就只能缩放成该尺寸检测精度和召回都会受损。设置动态轴后推理时传入任意尺寸ONNX Runtime会为每次请求做临时的shape推断代价是首次执行有可感知的开销可在服务启动时预热一张图。导出后建议做一次输出对比分别用Pytorch和ONNX Runtime喂同一输入比较输出张量的最大绝对误差。误差大于1e-3优先检查opset版本和BatchNorm状态model.eval()漏写会让BN的统计量在推理阶段继续漂移是OCR模型在测试和部署阶段表现不一致的头号原因。map_locationcpu这个参数也值得单独讲如果原来checkpoint存在GPU上导出时不在CPU上做一次load后续部署到无GPU机器就会报设备不匹配的错误。导出识别模型同理只是在导出前要把vocab.json一起打包部署侧用同一个文件做解码避免训练和推理词表不一致导致的乱码。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进