资讯动态

PyTorch+OCR实现火车车厢号识别:从检测到部署

发布时间:2026/10/1 17:10:09 来源:尧图企业网站定制
简介面向铁路货运管理、物流追踪与智能交通等对识别效率与精度要求较高的应用场景这套基于PyTorch框架的OCR深度学习方案专门解决火车车厢编号的自动识别与提取问题。从图像批量预处理、文字区域检测到序列识别代码覆盖了完整算法链路既可用于自动化录入车厢编号、清点查验车辆也能复用到车牌、集装箱号等类似识别任务中。压缩包共41个文件以25个Python脚本为核心涵盖模型定义、训练与测试入口、配置参数、工具函数等多个模块另有11个txt说明文档用于环境配置2个pkl标签映射文件保存字母与数字索引以及docx、json、md等配套材料补充设计原理与使用说明整个资源包仅86KB并划分了检测、识别、数据生成等子模块结构清晰。目前已有31人学习使用适合具备深度学习基础、希望深入理解OCR落地细节的工程师和研究者。结合附赠文档与测试示例可快速掌握CRNN、空间变换网络等组件的原理与调用方式明确从图像输入到编号输出的特征提取与解码逻辑同时参考验证码生成等数据增强手段缩短从零开始在真实铁路货运场景中部署调优的时间。1. 火车车厢号识别系统到底在解决什么问题从一列货车上千张图说起火车车厢号识别系统听上去是个窄项目但它是铁路货运管理里最刚需的“数字入口”。一辆货运列车几十节车厢每节车厢的编号是唯一的身份标识物流追踪、运费核算、智能交通调度全都先要正确拿到这个编号。过去场站靠人工抄号夜间、污损、反光都容易抄错效率也低。这个项目标题的落点就是用 PyTorch 搭一个 OCR 深度学习模型从抓拍图像里自动定位车厢号区域并提取字符串替代人工录入支撑自动化处理大量车厢图像。适合接这个项目的是做货运站信息化、物流园区闸口自动化和智能交通系统集成的工程师尤其是已经攒了一批车厢图像、正在为误识别和人工成本头疼的团队。2. 选型与数据为什么是 PyTorchOCR而不是传统模板匹配2.1 车厢号识别的难点污损、遮挡、反光与多角度先看清对象再选技术。车厢号是喷在侧板上的字符串通常由车型代码和数字编号组成但它的识别条件比车牌恶劣得多。车牌有相对统一的字体和尺寸车厢号则因车辆段不同、喷刷批次不同而字体各异同一列列车里新旧车厢混编字符清晰度也完全不同。货运车辆长期露天运行侧板上的污渍、雨水、铁锈会直接污染字符区域货场照明在夜间形成的反光带又会把字符局部吞掉。再加上相机安装位置很难正对车厢号图像往往带着透视和倾斜传统模板匹配在这种组合下基本没法通用。车厢号还不是图像里唯一的大字标记。车厢侧板上同时印着载重、容积、换长、路徽等信息OCR 如果整张图直接识别很容易把旁边这些标记读进来。这个项目必须先解决“车厢号在哪里”的问题再解决“这段字符是什么”的问题所以技术路线天然是两阶段目标检测定位车厢号区域再做文本识别。2.2 技术选型理由PyTorch与OCR模型CRNN/CTC、检测识别两阶段选 PyTorch 不是因为热门而是因为可控。OCR 识别车厢号需要自定义字符集、输入尺寸、后处理规则还要在货运站现场做离线部署不能每次识别都请求外部接口。PyTorch 深度学习模型可以完整打通从数据加载到 ONNX 导出的链路调试时也能把中间特征图拉出来看这在 Commons 商用 OCR 服务里很难做到。PaddleOCR 这类现成框架也不是不能用但面对“车厢号车型代码”这种强格式场景定制字符集和后处理的成本有时候比从零训练一个识别头还高。两阶段方案里检测部分我用轻量目标检测模型框出车厢号区域识别部分用 CRNNCTC。CRNN 的意思是卷积网络提取视觉特征循环网络对特征序列建模最后用 CTC 对齐不定长字符序列。它最实在的好处是不需要逐字符标注标注车厢图像时只要给一个框和整串车厢号就能监督训练。也有人用 Transformer-based 的识别网络精度上限可能更高但参数量和数据需求上去了在货场那种只有一块普通 GPU 的机房环境里CRNN 更容易跑出一个稳定可用的模型。这个章节背后还有一个选型细节PyTorch 环境搭建不是一句话能带过的驱动、CUDA、PyTorch 版本三者要匹配后面避坑章节我会专门讲。先把数据做扎实否则后面训练阶段全在跟“数据太脏”较劲。2.3 数据准备车厢号图像采集与标注建议数据是决定这个项目能不能落地的第一道关。采集车厢号图像不要只从视频监控里截帧那样视角太单一。我建议在车厢两侧不同高度、不同俯仰角、不同光照下各拍一批同一种车型至少保证两组拍摄机位。标注时不需要给每个字符画框只要框出车厢号整体区域标签填整串文本字符切分交给 CTC 去学。VOC 和 COCO 格式都常见但为了后面训练方便我习惯先把标注转成 CSV。下面是一个把 VOC XML 转成 CSV 的脚本import os import xml.etree.ElementTree as ET import csv def voc_to_csv(anno_dir, image_dir, output_csv): rows [] for xml_file in os.listdir(anno_dir): if not xml_file.endswith(.xml): continue tree ET.parse(os.path.join(anno_dir, xml_file)) root tree.getroot() filename root.find(filename).text image_path os.path.join(image_dir, filename) for obj in root.findall(object): label obj.find(name).text box obj.find(bndbox) xmin int(float(box.find(xmin).text)) ymin int(float(box.find(ymin).text)) xmax int(float(box.find(xmax).text)) ymax int(float(box.find(ymax).text)) rows.append([image_path, label, xmin, ymin, xmax, ymax]) with open(output_csv, w, newline, encodingutf-8) as f: writer csv.writer(f) writer.writerow([image_path, label, xmin, ymin, xmax, ymax]) writer.writerows(rows) if __name__ __main__: voc_to_csv(annotations/, images/, cargo_bbox.csv)这个脚本把标注 XML 里的文件名、目标名称和检测框四坐标读取出来落成一行 CSV。参数上要注意anno_dir和image_dir的目录结构必须对应图片文件名里如果有中文或空格先统一重命名否则后续 PyTorch DataLoader 读路径时很容易遇到编解码问题。标签字段一定要保存完整车厢号字符串不要拆成单个字符。数据整理完先做一轮清洗模糊、严重过曝、车厢号被遮挡超过三分之一的图片直接删掉否则模型会花大量容量去拟合噪声而不是字符特征。识别模型的输入高度我固定在 32 像素宽度按比例缩放后补齐到 160 或 200。字符集提前枚举包含数字 0-9、大写字母 A-Z以及车厢号里可能出现的连接符。把这个字符集固定下来训练和推理共用同一份映射表后面就不会出现“训练时没见过这个字符”的尴尬情况。3. 用PyTorch实现车厢号识别从检测到文本识别的最小方案3.1 第一阶段目标检测定位车厢号区域车厢号在图像里的位置不固定尤其列车连接处车厢号可能紧贴车厢边框。检测阶段输出的框要尽量贴住车厢号不能把“载重”“换长”这些邻近标记圈进来。完整训练一个检测网络代码很长常见做法是直接用 YOLO 系列做检测。但无论用哪种检测头数据装载这一步绕不开下面这个 Dataset 封装是 PyTorch 训练检测任务的常见写法import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T class BoxDataset(Dataset): def __init__(self, csv_file, img_size640, is_trainTrue): self.samples [] with open(csv_file, r, encodingutf-8) as f: for line in f.read().strip().splitlines()[1:]: img_path, label, xmin, ymin, xmax, ymax line.split(,) self.samples.append((img_path, label, int(xmin), int(ymin), int(xmax), int(ymax))) self.img_size img_size self.is_train is_train def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label, xmin, ymin, xmax, ymax self.samples[idx] img Image.open(img_path).convert(RGB) w, h img.size scale self.img_size / max(w, h) new_w, new_h int(w * scale), int(h * scale) img T.Resize((new_h, new_w))(img) pad_w (self.img_size - new_w) // 2 pad_h (self.img_size - new_h) // 2 img T.Pad((pad_w, pad_h, self.img_size - new_w - pad_w, self.img_size - new_h - pad_h), fill0)(img) xmin_n (xmin * scale pad_w) / self.img_size ymin_n (ymin * scale pad_h) / self.img_size xmax_n (xmax * scale pad_w) / self.img_size ymax_n (ymax * scale pad_h) / self.img_size boxes torch.tensor([[xmin_n, ymin_n, xmax_n, ymax_n]], dtypetorch.float32) labels torch.tensor([1], dtypetorch.long) img T.ToTensor()(img) return img, boxes, labels这段代码的关键是坐标换算图像先做等比缩放到最长边 640 再居中填充缩放系数和填充边距必须同步换算到检测框坐标上否则框位置和图像内容对不上。labels这里用单类别 1 代表车厢号如果检测任务同时还要识别车型代码、载重标识等区域就要维护一个类别映射表。检测模型本身可以用 YOLO 或 Faster R-CNN 训练但推理时有一个容易忽略的细节检测框不能直接裁给 OCR宽度建议外扩 5%高度外扩 10%不然车厢号最后一个字符很容易被截断。3.2 第二阶段OCR文本识别模型CRNNCTC 或 预训练OCR识别模型我用 CRNN它的结构比较固定卷积层提取特征双向 LSTM 建模序列线性层输出每个时间步的字符概率。这里给一个精简实现import torch.nn as nn import torch class CRNN(nn.Module): def __init__(self, num_classes, height32, width160): super().__init__() self.cnn nn.Sequential( nn.Conv2d(3, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), nn.Conv2d(128, 256, kernel_size3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.MaxPool2d((2, 1), (2, 1)), ) self.rnn nn.LSTM(256 * (height // 8), 128, bidirectionalTrue, num_layers2, batch_firstTrue) self.fc nn.Linear(256, num_classes) def forward(self, x): x self.cnn(x) b, c, h, w x.size() x x.reshape(b, c * h, w) x x.permute(0, 2, 1) out, _ self.rnn(x) out self.fc(out) return out这个模型的输入是(batch, 3, 32, 160)经过三次池化后特征图高度变成 4再 reshape 成宽度方向的序列交给双向 LSTM。MaxPool2d((2, 1), (2, 1))只在高度上降采样宽度不变这样特征序列长度和输入宽度是对应的。训练时num_classes是有效字符数加 1多出来的一个位置给 CTC 的 blank。实际项目里我会换用预训练 CNN 作为骨干比如 ResNet18 的前几层能明显加速收敛但这个最小结构读起来更容易理解模型在做什么。3.3 训练配置与参数批次、学习率、图像尺寸、字符集、增强训练循环本身不长但参数设置直接决定模型最终能不能用。以下代码是一套能跑通的最小训练流程import torch import torch.nn as nn from torch.utils.data import DataLoader # charset 里不包含 blankblank 由 CTCLoss 单独指定索引 charset list(0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ-) char2idx {c: i for i, c in enumerate(charset)} idx2char {i: c for c, i in char2idx.items()} def collate_fn(batch): images, labels zip(*batch) images torch.stack(images) targets [torch.tensor([char2idx[c] for c in lb], dtypetorch.long) for lb in labels] target_lengths torch.tensor([len(t) for t in targets], dtypetorch.long) return images, targets, target_lengths train_loader DataLoader(train_set, batch_size32, shuffleTrue, collate_fncollate_fn) model CRNN(num_classeslen(charset) 1) optimizer torch.optim.AdamW(model.parameters(), lr1e-4) criterion nn.CTCLoss(blanklen(charset)) device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) for epoch in range(80): model.train() for images, targets, target_lengths in train_loader: images images.to(device) logits model(images) log_probs logits.log_softmax(2).permute(1, 0, 2) input_lengths torch.full((images.size(0),), logits.size(1), dtypetorch.long) loss criterion(log_probs, targets, input_lengths, target_lengths) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() if epoch % 10 0: print(fepoch {epoch}, loss {loss.item():.4f})代码里CTCLoss的 blank 索引设为len(charset)所以模型输出维度是字符数加一。input_lengths是模型输出的序列长度CTC 要求它不比真实标签短target_lengths是每个样本车厢号的字符个数。梯度裁剪到 5.0 是防止 LSTM 在训练后期梯度爆炸这个值来自经验设太小收敛慢设太大容易出现 loss 跳变。图像尺寸方面输入高度 32、宽度 160 比较省显存如果车厢号字符数多或者字形宽建议宽度加到 240。增强不要做随机大角度旋转车厢号本身是水平的转多了反而引入偏离真实分布的样本。常用增强是随机亮度扰动、少量高斯噪声、随机透视变换模拟拍摄机位以及用随机光斑模拟反光。批大小在 16 到 32 之间学习率从 1e-4 开始到后 20 个 epoch 降到 1e-5。4. 部署与调优把模型跑进货运现场的约束4.1 推理流程图像预处理、检测后处理、OCR后处理正则、校验模型训练完部署阶段要处理的坑比训练还多。下面是一段完整的单图推理流程包括检测、裁剪、识别和 CTC 解码from PIL import Image import torch import torchvision.transforms as T def infer_car_number(img_path, detect_model, ocr_model, transform, devicecpu): img Image.open(img_path).convert(RGB) img_tensor T.ToTensor()(img).unsqueeze(0).to(device) with torch.no_grad(): pred detect_model(img_tensor)[0] boxes pred[boxes][pred[scores] 0.5] w, h img.size result [] for box in boxes: x1, y1, x2, y2 [int(v) for v in box.tolist()] pad_y int((y2 - y1) * 0.1) pad_x int((x2 - x1) * 0.05) crop img.crop((x1 - pad_x, y1 - pad_y, x2 pad_x, y2 pad_y)) crop crop.resize((160, 32)) crop_tensor transform(crop).unsqueeze(0).to(device) with torch.no_grad(): logits ocr_model(crop_tensor) pred_ids logits.argmax(dim-1).squeeze(0).tolist() chars [] prev None for idx in pred_ids: if idx ! blank_idx and idx ! prev: chars.append(idx2char[idx]) prev idx result.append(.join(chars)) return resultpad_x和pad_y是经验值分别按检测框宽高的 5% 和 10% 向外扩避免字符被裁边。scores 0.5是阈值现场误检多就调到 0.7漏检多就往下调更稳的做法是不设置单一阈值先留下所有大于 0.1 的框再用后处理规则过滤。CTC 解码必须做“连续相同字符合并再跳 blank”这个逻辑如果漏了识别结果会把“1122”读成“12”这是一个非常容易翻车的细节。部署阶段的另一个关键点是后处理校验。车厢号有固定格式常见的是车型代码加数字编号车型代码可以枚举出来。识别结果如果不在格式范围内系统宁可丢弃也不要提交入库这一点比模型准确率更影响现场口碑。4.2 提速与压缩PyTorch转ONNX、TensorRT、批处理PyTorch 模型直接部署到现场服务里太重我一般先导出 ONNX再用 ONNX Runtime 或 TensorRT 做推理加速。导出代码很简单import torch from model import CRNN model CRNN(num_classeslen(charset) 1) checkpoint torch.load(best_model.pth, map_locationcpu) model.load_state_dict(checkpoint[state_dict]) model.eval() dummy_input torch.randn(1, 3, 32, 160, devicecpu) torch.onnx.export( model, dummy_input, ocr_model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch}, output: {0: batch}}, opset_version12 ) print(ONNX导出完成)导出 ONNX 的关键是动态轴。dynamic_axes把 batch 维度设为动态可以在服务端一次推理多张图但如果转换时报 LSTM 算子不兼容就把 opset_version 降到 11或者去掉动态轴改用固定 batch1。TensorRT 不是所有算子都能跑满加速双向 LSTM 在部分 TensorRT 版本里优化效果不稳定所以我的建议是先接 ONNX Runtime 保证功能正确再试 TensorRT 的 FP16确认精度不掉再切换。如果现场要处理整列车的图像单张推理太慢。常见做法是把同一节车厢连续抓拍的多帧图像攒成一个 batch一次推理多张吞吐量能提升好几倍。批量推理时图片宽度不一致要么提前 pad 成固定宽度要么使用动态输入尺寸否则 batch 里的 Tensor 形状对不上会直接报错。4.3 现场鲁棒性调优多帧投票、置信度阈值、反光/雨雾图像增强现场识别最怕的不是单帧准确率低而是单帧输出一个错误但看起来很像真的车厢号。要压住这种风险我通常用多帧投票from collections import Counter def vote_result(frame_results, min_conf0.6, min_frames2): valid [r for r in frame_results if r[conf] min_conf] if len(valid) min_frames: return None, 0.0 texts [r[text] for r in valid] counter Counter(texts) text, count counter.most_common(1)[0] return text, count / len(valid)这个投票不是简单取出现次数最多的字符串更稳的做法是给每一帧的置信度做加权。如果同一个车厢号出现了三帧其中两帧置信度 0.9 识别为“C64K 1234567”一帧置信度 0.55 识别成“C64K 1234561”加权后正确答案仍然会胜出。现场安装时让相机与车厢运行方向呈 45 度角能拍到更多角度的车厢号投票效果也会更好。夜间反光问题很难在推理阶段彻底解决我的习惯是在训练阶段加“反光模拟”增强把正常车厢号图像随机叠加一块高亮光斑让模型学会忽略那种一片白的区域。这类增强 OpenCV 就能做不依赖外部库也比部署后再想“去反光算法”要直接得多。雨雾天气同理训练数据里混入低对比度样本比上线后临时调阈值靠谱。5. 避坑车厢号识别项目里我踩过的5个坑5.1 现象夜间图像识别率骤降白天正常夜间车厢号照片大多靠货场补光灯补光灯在侧板上形成一条高亮光带字符局部被光斑吞掉。而训练数据里白天图像占比太高导致模型没有见过这种极端反光分布。解决采集一批夜间实拍图按不同曝光强度回灌训练集同时用随机光斑增强模拟高光。推理时不要拿单帧去赌必须配多帧投票一帧反光看不清相邻两帧很可能能补上。5.2 现象0 被识别成 O8 被识别成 B模型在数字和字母之间反复横跳车厢号字符集同时包含数字和字母0 和 O、8 和 B 这种相似字形本来就是 OCR 的经典难题。CTC 模型只学习字符上下文不知道整个车厢号必须满足“车型代码数字编号”的强格式。解决在识别结果后面加格式校验和纠错模块。先把可能出现的车型代码枚举出来识别结果先拆成车型段和数字段再按规则纠正明显不合理的字符。也可以训练一个 n-gram 语言模型给候选字符串打分这是传统 OCR 工程里比模型本身更值得花时间的部分。5.3 现象PyTorch环境搭建时训练报错“CUDA error: no kernel image is available”这个问题我在 PyTorch 环境搭建时踩过现象是一跑model.to(device)就报错或者训练第一个 batch 直接崩。原因是显卡驱动支持的 CUDA 版本和 PyTorch 自带的 CUDA runtime 不匹配常见于新显卡配旧版 PyTorch或者在 WSL2 里安装了和宿主机驱动不对应的 CUDA 组件。解决先运行nvidia-smi看驱动支持的最高 CUDA 版本再按这个版本选对应的 PyTorch 安装命令。WSL2 里不需要单独装显卡驱动直接用宿主机映射进来的驱动就行。装完先跑两行代码确认环境import torch print(torch.__version__) print(torch.cuda.is_available())如果输出 True再开始训练否则后边所有报错都可能是环境问题造成的。5.4 现象训练 loss 降得很低但验证集识别准确率一直上不去这种情况通常是模型学到的是背景噪声不是车厢号字符。货运车厢图像里侧板纹路、锈迹、旁边的载重标记都和字符区域混在一起如果数据量少模型很容易记住这些干扰模式。解决识别模型不要直接拿整张车厢图训练一定先做检测裁剪只把车厢号区域喂给 CRNN。训练时加入背景裁剪作为负样本让模型学会区分“这是车厢号区域”和“这不是”。如果数据量低于 1 万张先冻结 CNN 骨干只训练 LSTM 和分类层等数据攒够了再解冻整体微调。5.5 现象TensorRT 转换后速度提升了但某些帧的识别结果和 PyTorch 不一致这不是模型代码问题是精度问题。TensorRT 在算子融合和低精度推理时会改变中间层数值精度双向 LSTM 在某些 TensorRT 版本里支持不完整INT8 量化后更是如此。解决先转 FP16用同一批验证图对比 PyTorch 和 ONNX Runtime 的输出字符准确率掉得超过 0.5% 就要查具体差异帧。不要为了速度直接上 INT8除非现场所有车厢号都有双通道比对兜底。保留一份 PyTorch CPU 权重作为后台对关键车厢号做二次校验不一致时以高置信度结果为准。6. 验证与进阶让识别系统从“能跑”到“可信”6.1 用字符级指标衡量模型而不是“识别率”项目验收时不要只看“准确率 97%”这种一句话。车厢号识别要拆成两个指标字符级准确率和整串准确率。字符级准确率是预测字符串和目标字符串逐字符对比的正确率整串准确率要求完全一致。对货运管理来说整串一致才算一次正确识别。def char_accuracy(pred, target): pred_chars list(pred) target_chars list(target) matches sum(1 for a, b in zip(pred_chars, target_chars) if a b) return matches / len(target_chars) if target_chars else 0.0验证集划分也要讲究。同一个车厢的连续抓拍帧只能放在同一组里不能训练集和验证集里都出现同一节车厢的不同角度照片否则指标会虚高。我一般会把数据按“拍摄日期”切分用前两周的数据训练后一周的数据验证这才能反映模型在没见过的新车厢上的表现。6.2 半监督自训练与持续迭代现场跑起来之后模型会积累大量新的车厢图像。与其等人再标注一轮不如用半监督自训练把置信度高于 0.95 的识别结果当作伪标签加入训练集重新训练。这个阈值要卡得比较高哪怕只增加几百张高质量伪标签样本也比随机采样的效果稳定。多摄像头场景下可以把两路相机对同一节车厢的识别结果做交叉确认两边一致才进入伪标签池。持续迭代的关键是建立一个难例回传机制。现场每次出现低置信度识别、或者被后处理规则判为格式错误的样本都单独存下来定期人工复核后再补进训练集。OCR 这种黑匣子模型越往后越靠数据闭环而不是网络结构调参来提精度。我现在接到类似项目会先花两天把评估集和后处理搭出来再开始训练模型。这个顺序倒过来后面大概率要返工。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑