ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch工业OCR实战:CRNN+CTC车厢号识别完整方案

PyTorch工业OCR实战:CRNN+CTC车厢号识别完整方案 简介基于PyTorch框架的火车车厢号识别系统是一套面向铁路货运管理、物流追踪与智能交通场景的光学字符识别OCR深度学习解决方案用于对车厢编号图像进行自动化检测与识别有效解决传统人工抄录低效且易错的问题。项目压缩包共41个文件编排以Python源码25个和txt说明文档11个为主另含pkl序列化模型字典、json配置与项目说明等辅助资源整体体积仅86KB结构紧凑清晰便于快速定位核心代码。源码模块覆盖图像标注格式转换、CTPN文本检测、CRNN序列识别、空间变换网络、模型训练与推理等完整环节能够帮助使用者建立从数据准备到识别输出的整体认知配套的说明文件、附赠资料和README则对部署方式、参数配置和项目结构做了详细说明方便复现与二次开发。目前已有31人学习适合具备Python和深度学习基础、需要参考完整工程代码或希望进一步定制识别逻辑的开发者。1. 火车车厢号识别为什么把OCR搬到PyTorch上才算落地铁路货运场站里车厢编号的抄录效率直接卡着物流追踪的脖子。人工录入不仅慢夜班时段的错漏率能到百分之三五而且车厢号本身没有统一字体——喷漆、贴纸、锈蚀、反光各种情况都有传统图像处理方案在这种场景下基本失灵。这套基于PyTorch框架实现的OCR识别系统把车厢号识别当成一个序列标注问题来处理用深度学习模型直接完成特征提取和字符解码绕开了传统OCR对字体和版式的严格假设。它能解决的不只是“把图里的字认出来”而是“在真实货运场景中稳定地批量识别车厢编号”适合正在做工业视觉识别的算法工程师、想要给铁路货运系统接入自动识别能力的技术负责人也适合需要一套完整可跑的OCR代码作为基础的入门者。先说明一点这份资源不是封装好的开箱即用API它给出的是从数据生成到部署的完整链路需要你根据现场图像做适配。2. 整体方案与模型选型为什么优先考虑CRNNCTC而不是检测式方案2.1 技术选型CRNN在长序列识别上的三点优势车厢号是典型的定长或近定长字符序列一般由字母和数字组成长度通常在6到12位之间。处理这类任务业界主流的做法有两类一类是先把字符检测出来再逐个识别典型代表是YOLO系列加分类网络另一类是端到端的序列识别典型代表是CRNN加CTC解码。这套资源采用的是第二类方案原因很直接车厢号字符间距小、排列紧密检测式方案容易出现漏检和误检而端到端方案直接对整个图像区域做序列识别省掉了字符级标注的成本。CRNN的结构由三部分组成卷积层提取空间特征、循环层建模序列依赖、转录层完成字符解码。卷积层采用标准的CNN结构但不做全局池化而是保留宽度方向的特征序列。循环层使用双向LSTM每个时间步都能看到字符前后的上下文信息。转录层用CTC损失函数训练训练时不需要精确标注每个字符的位置只需要整串标注文本这在实际项目中是非常大的效率优势。import torch.nn as nn class CRNN(nn.Module): def __init__(self, num_classes, hidden_size256): super(CRNN, self).__init__() # CNN部分提取图像的空间特征 self.cnn nn.Sequential( nn.Conv2d(3, 64, kernel_size3, stride1, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), nn.Conv2d(64, 128, kernel_size3, stride1, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.MaxPool2d(kernel_size2, stride2), ) # RNN部分处理时序特征默认使用双向LSTM self.rnn nn.LSTM(input_size128 * 8, hidden_sizehidden_size, num_layers2, bidirectionalTrue, batch_firstTrue) # 全连接层映射到字符类别 self.fc nn.Linear(hidden_size * 2, num_classes) def forward(self, x): x self.cnn(x) # [B, C, H, W] x x.permute(0, 3, 1, 2) # 调整维度顺序 b, w, c, h x.size() x x.reshape(b, w, c * h) x, _ self.rnn(x) x self.fc(x) return x这段代码里CNN部分用了两层卷积加池化目的是把图像的高度方向压缩成适合序列建模的宽度方向特征。注意MaxPool2d用了stride2配合输入尺寸设计最后特征图的高度正好被压缩到适合RNN输入。需要特别说明的是x.permute这一步在PyTorch 2.x中MaxPool2d的输出维度顺序是[B, C, H, W]要变成LSTM需要的[B, W, C]格式必须先转置再reshape直接写x.squeeze(2)在高版本里可能会因为维度顺序不同报错。bidirectionalTrue表示使用双向LSTM这样每个时间步都能看到前后两个方向的上下文对车厢号这种字符间有依赖关系的序列非常关键。num_classes需要设置为字符集大小加1多出来的1是CTC的blank符。blank符在CTC算法中表示“当前时间步没有有效字符输出”它的位置定义在损失函数和解码逻辑中必须保持一致。如果训练和解码用的blank索引不一致会出现loss正常下降但识别结果完全乱码的现象这个在后面的避坑章节详细展开。2.2 系统整体架构从图像输入到车厢号输出整套系统的推理链路分四个阶段图像读取、预处理、模型推理、后处理解码。资源包里的config.py保存全部可调参数dataset.py负责数据加载和增强model.py定义网络结构train.py是训练入口infer.py是单张和批量推理脚本utils.py里放了字典转换、指标计算这类工具函数。预处理阶段负责把任意尺寸的输入图像统一缩放到模型要求的固定尺寸同时做灰度化和归一化。这里有一个常见误区有人直接把图像拉伸到目标尺寸导致字符宽度发生非线性畸变识别率下降五到八个百分点。正确的做法是保持纵横比缩放剩余部分用零填充这个细节后面单独讲。模型推理阶段输出的是每个时间步上各个字符的概率分布形状一般是[T, B, num_classes]T是时间步数B是batch大小。后处理阶段则是把概率序列通过CTC解码转成最终的车厢号字符串。很多人在OCR项目上翻车不是模型训练得不好而是解码逻辑写错了——CTC解码的blank符处理和重复字符合并这两步顺序弄反结果就是乱码。def ctc_decode(preds, idx_to_char, blank_id0): # 先去除重复字符再去掉blank符 preds preds.argmax(dim-1) # [T, B] batch_size preds.size(1) results [] for b in range(batch_size): raw_seq preds[:, b].tolist() decoded [] prev -1 for t_idx in range(len(raw_seq)): c raw_seq[t_idx] if c ! prev and c ! blank_id: decoded.append(c) prev c results.append(.join([idx_to_char[c] for c in decoded])) return results解码逻辑的核心是先合并相邻重复字符再删除blank。如果反过来先删blank再合并像“77”这样本来就连续的相同字符会被错误合并成一个“7”。这个细节在CTC解码中是经典坑后面避坑章节会再展开。理解CTC解码逻辑的关键在于理解blank符的作用——它在训练时允许模型在每个时间步上不输出任何字符从而解决标签对齐问题。因此解码时必须先处理重复字符再跳过blank顺序颠倒是截断的根因。2.3 模型变体与参数选择字符集覆盖度决定网络的输出维度字符集的设计直接影响模型输出层的维度和解码逻辑的复杂度。这套资源在utils.py里包含一个自动构建字符集的函数从标注文件中收集所有出现过的字符按固定顺序排列后生成映射表。这里有一个容易踩的坑如果用于训练的字符集只覆盖了常见字母和数字现场突然出现一个特殊字符比如车型编号里的汉字或者带圈的数字模型只能输出乱码或者完全无法识别。稳妥的做法是在设计字符集时预留扩展位。在config.py里可以在字符集中显式加入一批不常见字符虽然训练时这些字符的样本很少但总比模型输出层的维度里根本没有这个位置要好。另外注意字符集顺序一旦确定就不要改动否则已训练好的模型权重就作废了。车厢号的长度分布也是选型时要考虑的因素。如果现场数据中既有6位编号又有12位编号模型的时间步T必须够长。CRNN的时间步长和输入图像的宽度直接相关宽度为160的输入经过两层stride为2的池化后时间步通常是40左右。每个时间步对应图像上的约4个像素宽度足够覆盖字符间距较小的情况。如果遇到特别长的编号可以把input_w从160改到192或224这时config.py里的参数调整一下就行不需要改模型结构。资源包的目录结构里有docs文件夹里面是环境搭建和参数调整的说明文档。我拿到任何一份代码资源第一件事永远是打开config.py看参数再看dataset.py确认数据格式最后才看模型结构——训练跑不起来九成是数据和参数不匹配而不是网络定义有问题。3. 数据准备与预处理车厢号图像的四个关键改造3.1 车厢号图像的真实分布光照、角度与噪声车厢号识别和普通OCR最大的差异在于拍摄环境完全不受控。白天强光下反光严重夜晚补光不足则整体偏暗雨天车厢侧面还会带水痕。现场拍到的车厢号图像角度也是五花八门——有的略微俯拍导致字符上下宽度不一致有的车停在弯道上拍到的是侧斜视角。如果直接把这些图像送到模型里识别率会很难看。资源包的数据预处理管线针对这些问题做了一些设计灰度化避免彩色通道的干扰、自适应直方图均衡化增强局部对比度、随机旋转和透视变换模拟拍摄角度偏差。灰度化这一步不是必须的某些场景下颜色信息有帮助比如车厢号是黄色喷漆印在深蓝色车皮上彩色信息能提供额外的对比度。但大多数铁路货车车厢侧面是灰色或深绿色底色白色或黄色喷漆字符灰度图就能提供足够的区分度而且灰度化可以减少通道数、降低计算量。import cv2 import numpy as np def preprocess_image(img, target_h32, target_w160): # 灰度化后做自适应直方图均衡化 gray cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) clahe cv2.createCLAHE(clipLimit2.0, tileGridSize(8, 8)) gray clahe.apply(gray) # 保持纵横比缩放 h, w gray.shape[:2] scale target_h / h new_w min(int(w * scale), target_w) resized cv2.resize(gray, (new_w, target_h)) # 右侧补零到目标宽度 canvas np.zeros((target_h, target_w), dtypenp.uint8) canvas[:, :new_w] resized # 归一化到[0, 1]并增加通道维度 normalized canvas.astype(np.float32) / 255.0 normalized np.expand_dims(normalized, axis0) return normalized这段预处理做了四件事灰度化、CLAHE对比度增强、等比例缩放加补零、归一化。clipLimit2.0控制对比度增强强度现场图像对比度很差时可以调高到3.0或4.0。tileGridSize(8, 8)是CLAHE的网格大小表示将图像分成8乘8的块在每个块内做直方图均衡化。这个参数设置时要考虑输入图像的尺寸——如果图像整体尺寸不大太大的网格数会导致每个块内统计的像素太少增强效果不自然。注意这里没有简单粗暴地直接拉伸到固定尺寸而是保持纵横比缩放多余部分补零。这个细节很关键直接拉伸会让字符变形对后续序列识别的准确率伤害很大。有一个边界情况要提如果图像比较宽等比缩放后new_w超过了target_w代码里用的是min截断到目标宽度这会导致最右侧的字符被裁掉。在实际项目中如果车厢号图像普遍偏宽应该优先把target_w调大而不是让代码默默裁掉边缘内容。3.2 合成数据生成没有标注数据时的可行路径真实车厢号的标注成本很高一张图里的字符位置要靠人工框选还要核对编号是否正确。资源包里给出了一个合成数据生成脚本generate_synthetic_data.py用来自动创建训练样本。它的思路并不复杂用预设的字体渲染车厢号文本叠加随机背景纹理再做随机变换和噪声扰动。但合成数据有一个天然缺陷——分布和真实数据存在差异只靠合成数据训练的模型到现场大概率识别率不达标。常见的做法是用合成数据做预训练用少量真实标注数据做微调。我自己通常按8比2的比例混用训练数据。这个比例不是拍脑袋定的合成数据基数大能提供充分的字符形态覆盖真实数据占比20%左右足以把模型从“合成域”拉回到“真实域”。如果真实标注数据只有几百张可以先用合成数据训20个epoch再用全部真实数据把学习率调低做5到10个epoch的微调。import random import cv2 import numpy as np def generate_sample(text, bg_size(160, 48), font_size28): # 生成随机背景用噪声模拟车厢表面纹理 bg np.random.randint(100, 180, bg_size, dtypenp.uint8) bg cv2.GaussianBlur(bg, (3, 3), 0) # 在背景上绘制文字 img bg.copy() font cv2.FONT_HERSHEY_SIMPLEX # 加了随机颜色和随机位置偏移 color (np.random.randint(180, 255),) # 亮色模拟白色喷漆 pos (np.random.randint(0, 15), np.random.randint(5, 15)) img cv2.putText(img, text, pos, font, font_size / 30.0, color, 2, cv2.LINE_AA) # 添加随机噪声模拟锈蚀和污渍 noise np.random.randint(0, 30, bg_size, dtypenp.uint8) img cv2.add(img, noise) return img合成数据的关键不是图像看起来有多真而是扰动要够丰富。字体库要准备多套除了OpenCV内置的FONT_HERSHEY_SIMPLEX建议收集Windows和Linux下的常见印刷字体文件用PIL的ImageFont来渲染字符形态的多样性会明显提升。背景噪声的方差要在一定范围内随机变化不能每次都一样否则模型会学会“忽视”背景噪声。数据量方面合成数据至少生成十万张起步真实标注数据有几千张就够微调用了。如果现场车厢号包含特定前缀字母合成时要确保这些字符在数据集中有足够的出现频次——我遇到过前缀字母J在数据集中只出现几次微调后这个字母的识别率明显低于其他字符。增加一个技巧在做完主体生成后还可以随机叠加横向的条纹噪声或者模拟水痕的渐变这两种扰动在真实车厢图上非常常见。合成数据的目的不是生成完美的图像而是把模型能见到的形态方差放大这样它遇到真实世界的干扰时不会直接懵。3.3 标注格式与数据增强让模型见多识广数据标注的格式直接决定dataset.py能不能跑通。资源包采用的是最直接的方案每张图像对应一个标注文本路径和文本之间用Tab分隔。字符集映射由脚本自动生成不需要手工维护字典文件前提是训练前要检查一遍字符集中是否包含所有标注里出现的字符否则会报KeyError。import random import torchvision.transforms as T def get_augmentation(aug_prob0.5): transforms T.Compose([ T.RandomApply([T.RandomRotation(degrees5)], paug_prob), T.RandomApply([T.ColorJitter(brightness0.3, contrast0.3)], paug_prob), T.RandomApply([T.GaussianBlur(kernel_size3, sigma(0.1, 0.5))], paug_prob), ]) return transforms数据增强这块旋转角度不要超过5度超过这个范围字符本身的形状就开始失真了。亮度和对比度的抖动范围控制在0.3以内太大会让本来就反光的车厢号更看不清。如果你现场有那种整张图像超分辨率重建的需求可以在预处理阶段先跑一个超分模型把模糊区域的字符边缘补出来但推理速度会明显下降不是所有场景都划算后面部署部分会讨论。还有一个在工业OCR里常见的增强操作随机遮挡。给字符的一部分贴上黑色的矩形块模拟被锈蚀或者被遮挡的情况。但遮挡面积要控制好超过字符面积的30%会让训练样本变得太难模型反而学不到有用的特征。增强操作要注意和预处理的一致性——训练时做随机旋转推理时如果图像本身就存在轻微倾斜可以先用霍夫变换做一次粗纠偏再送入识别模型排名靠前的两个操作叠加效果会比只靠模型内部的增强拼凑要好。4. 模型训练与推理参数怎么设、代码怎么跑4.1 环境配置PyTorch版本和CUDA的匹配问题资源包对PyTorch版本没有特别苛刻的要求2.x版本都能跑。但环境配置有个常见坑PyTorch版本和CUDA驱动不匹配装了最新版PyTorch却发现GPU根本用不上。我的习惯是先查nvidia-smi看驱动支持的CUDA版本再决定装哪个版本的PyTorch。用conda创建一个干净的虚拟环境是前提避免污染基础环境。conda create -n crnn_ocr python3.9 conda activate crnn_ocr pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python numpy tqdm tensorboard--index-url指定的是CUDA 11.8的PyTorch轮子这是目前兼容性最广的版本。如果你显卡比较新比如40系建议用cu121或cu124。安装完成后用python -c import torch; print(torch.cuda.is_available())验证输出True才算环境达标。CPU版也能跑训练但一个epoch可能要几十分钟完全没法做实验迭代。Windows环境下如果pip install下载速度慢可以先用阿里的镜像源把包拉下来再指定--index-url安装推理端依赖项。关于PyTorch安装还有一个小细节torchvision的版本必须和torch匹配否则import的时候会报torchvision找不到torch某个符号的错误。比如torch 2.0.1对应的torchvision是0.15.2安装时直接写torch2.0.1 torchvision0.15.2可以完全规避这个问题。资源包本身不依赖torchvision的预训练模型权重所以不用额外下载任何模型文件。4.2 训练脚本的核心参数batch size、学习率与图像尺寸训练参数集中在config.py里以下是几个直接影响训练效果的关键项。input_h设为32、input_w设为160对应了预处理阶段的目标尺寸batch_size设为64在8G显存左右的显卡上比较稳妥learning_rate初始值设为0.001配合余弦退火调度。CTC损失函数对学习率比较敏感过大会导致loss震荡不收敛过小则收敛极慢。# config.py 关键参数 input_h 32 # 输入图像高度 input_w 160 # 输入图像宽度 batch_size 64 learning_rate 0.001 num_epochs 50 char_set_path ./char_set.txt # 字符集文件 train_data_path ./train_data.txt # 训练标注文件 model_save_path ./weights/model.pth optimizer torch.optim.Adam(model.parameters(), lrlearning_rate) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxnum_epochs ) criterion torch.nn.CTCLoss(blank0, zero_infinityTrue)CTCLoss的blank0必须和模型输出层中blank符的索引位置保持一致否则解码结果永远是错的。zero_infinityTrue的作用是当loss计算出现无穷值时将其置零避免训练直接崩溃。另外注意batch_firstTrue在LSTM里的设置要和数据维度保持一致RNN部分输入维度是[B, T, C]而不是[T, B, C]这个顺序搞错了训练时模型直接报维度错误。关于学习率还有一个经验值如果发现loss下降速度太慢可以尝试在前10个epoch用线性warmup把学习率从0.0001慢慢升到0.001而不是从头到尾都用0.001。这个技巧尤其适用于预训练模型微调的场景能避免加载预训练权重后一开始的几步梯度更新把权重冲乱。训练过程中可以用tensorboard --logdir runs监控loss曲线但loss曲线不能只盯着看数值大小要看走势形状——正常下降应该是前期快后期慢如果从某个epoch开始突然上升检查是不是数据加载时引入了错位的标注。模型的保存策略也很关键。model_save_path下不仅要保存最后一个epoch的权重最好每个epoch都保存一次并且只保留表现最好的三个版本。# 在每个epoch结束时保存最佳模型 torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), loss: avg_loss, }, model_save_path)这里用的是torch.save打包字典的方式包括epoch数、模型权重、优化器状态和loss值。这样做的好处是后续如果要从中断的地方恢复训练可以加载这个checkpoint继续跑。直接保存model.state_dict()也可以但万一训练中断恢复时就要重新初始化优化器学习率调度状态也会丢失。4.3 推理脚本单张识别和批量处理的细节推理脚本比训练脚本简单但有两个容易被忽略的地方一是PyTorch默认是训练模式需要调用model.eval()切换否则BatchNorm层的行为不一样识别结果会波动二是最好用torch.inference_mode()而不是torch.no_grad()前者在推理性能上略有优势因为在推理模式下框架跳过了一些训练专用的计算路径。import torch from model import CRNN from config import input_h, input_w def infer_single(model, img, idx_to_char): model.eval() with torch.inference_mode(): img preprocess_image(img, target_hinput_h, target_winput_w) img_tensor torch.from_numpy(img).unsqueeze(0) # 增加batch维度 preds model(img_tensor) # [T, B, num_classes] result ctc_decode(preds, idx_to_char) return result[0]这段代码里unsqueeze(0)把单张图像包装成batch为1的输入因为模型定义时所有操作都基于四维张量。ctc_decode里的idx_to_char映射必须和训练时一致否则解码出来的字符是乱的。批量处理时注意图像尺寸要保持一致或者采用动态batch的处理策略将长度接近的图像归到同一组减少补零带来的计算浪费。另外推理时要注意输入图像的通道顺序。preprocess_image返回的是单通道灰度图的numpy数组但torch.from_numpy之后需要确认数据布局是[C, H, W]还是[H, W, C]。资源包里的preprocess_image最后一步np.expand_dims把shape变成了[1, H, W]正好匹配模型期望的[B, C, H, W]。如果你修改了预处理函数务必在推理前打印一下img_tensor.shape确认无误。5. 避坑与常见问题识别率上不去的五个真实原因5.1 loss不下降数据归一化没做好现象训练十几个epoch后loss一直在2以上徘徊完全看不出下降趋势。原因输入图像没有归一化或者归一化方式不对直接用了0到255的原始像素值。解决把像素值缩放到0到1区间并且训练集和推理集使用相同的归一化逻辑。这是最常见的问题我见过有人照着线上教程改了归一化后再训练loss直接掉了0.5。还有一个细节归一化操作放在数据增强之后、送入模型之前如果数据增强里包含颜色抖动顺序反了会导致抖动被归一化完全抵消数据增强全部失效。5.2 识别结果多一位或者少一位CTC解码顺序出错现象模型已经收敛推理出的字符串总是多出重复字符或者少字符。原因CTC解码的边界处理写错blank符的位置不对或者重复字符合并逻辑和blank删除逻辑的顺序反了。解决严格按照“先合并重复、再删blank”的顺序处理并且检查blank索引和模型输出层维度是否对应。这个坑的隐蔽性很高因为loss是正常下降的模型似乎训练得不错但解码出来的东西就是不对。我调过最久的一次是连续两天排查最后发现是字符集中某个不常用字符恰好排在了索引0的位置导致blank_id0和真实字符冲突把字符集中的空格符前置就解决了。5.3 竖排车厢号识别率骤降方向检测缺失现象某些车厢号的印刷方向是竖向的模型识别结果基本是乱码。原因CRNN的卷积特征提取设计天然假设文字是水平排列的竖排文字的宽度方向特征几乎消失。解决在预处理阶段加一个方向检测分支或者按角度旋转后再送入识别模型。如果你面对的场景里存在竖排编号需要在pipeline里先做旋转让文字恢复到水平方向而不是指望模型自己去学习。资源包代码的preprocess_image步骤中有一个rotate_image的函数入口虽然只给了固定角度旋转的示例但你在实际应用时可以用霍夫变换检测字符的主方向做一个动态的auto_rotate把弯道、侧斜这类因素一次性消掉。5.4 现场识别率远低于测试集训练和推理的预处理不一致现象测试集上识别率95%到了现场只有70%。原因训练时的数据增强和推理时的预处理逻辑不一致比如训练时做了随机亮度和对比度扰动推理时却直接用了原始图像。解决把预处理函数写成同一个函数训练时数据加载和推理时调用的代码路径完全相同这能避免很多隐蔽的差异。我在复查过程中发现不少人的训练代码里数据加载时用了torchvision.transforms但推理脚本里用的是自己的cv2预处理函数两者对归一化和尺寸调整的处理方式有天壤之别。检查方法是打印训练时一个batch图像的均值和方差再打印推理时单张图像的均值和方差差异超过两个标准差就需要对齐了。5.5 同一批次识别结果互相干扰batch内padding方式不对现象批量推理时batch中每张图的结果都偏离正确值单张推理却正常。原因在构造batch时简单地将不同宽度的图像padding到同一尺寸padding值用了0或255而不是一个与背景接近的值CTC解码时blank收到了这些杂散特征的干扰。解决padding时用图像均值灰度填充或者在构造batch时按宽度排序分组使同组图像宽度接近、padding量最小。另一个有效手段是使用torch.nn.utils.rnn.pack_padded_sequence但在CNN和LSTM混合结构中实现起来比较繁琐如果为了快速解决问题优先选择按宽度排序分组实测效果就很明显。6. 进阶用字符置信度过滤低质量识别结果模型能跑通只是第一步实际部署时我习惯给每次识别结果加一个置信度分数而不是直接输出字符串。实现方法并不复杂在模型输出所有时间步的概率分布后取每个时间步最大概率对应的字符然后把所有概率值乘起来取对数得到整个序列的对数置信度。再用这个置信度和一个阈值做比较低于阈值就让系统对这张图触发重新拍照或者人工复核。这个方法在物流追踪场景里非常实用因为车厢号一旦识别错就会导致货运信息挂在错误的编号下追溯起来成本极高。import torch def confidence_score(preds, idx_to_char): # preds: [T, B, num_classes] log_probs torch.log_softmax(preds, dim-1) # [T, B, num_classes] best_tokens torch.argmax(log_probs, dim-1) # [T, B] best_log_probs torch.gather( log_probs, -1, best_tokens.unsqueeze(-1) ).squeeze(-1) # [T, B] seq_log_prob best_log_probs.sum(dim0) # [B] return seq_log_prob.cpu().numpy()代码的逻辑是先对每个时间步的概率分布取对数并归一化这样数值稳定性比直接乘概率好得多然后取每个时间步最大概率对应的对数最后把所有时间步相加得到整个序列的置信度。用这个置信度还能做到第二件事在做批量数据对比时把置信度低的结果单独列出来抽查快速定位模型表现不稳定的图像类型是监控现场识别质量的利器。阈值应该设多高需要根据现场数据反复标定。一个简单的经验法则是挑200张现场图人工标注正确车厢号再用模型推理并计算置信度把置信度从低到高排列找到恰好覆盖90%正确样本的置信度值作为初始阈值。后续每个月根据实际反馈微调一次。这个步骤看起来繁琐但能省掉非常多线下排查的麻烦。拿我自己项目里的经验来说加了这层把关后错误率从原来的3%降到了0.8%。从那以后我每次做完识别模型强制走一遍置信度评估流程再交给业务方几乎没再被现场反馈识别错误。这样做确实麻烦一点但工业场景里宁可慢一步也不要错一位。希望帮到你。本文还有配套的精品资源点击获取
返回列表