CNN交通标志分类实战:GTSRB数据预处理与PyTorch训练推理全流程
简介智慧交通系统中交通标志自动识别是视觉感知的重要环节这一基于卷积神经网络的分类项目实践包面向人工智能学习者、竞赛参与者以及需要完成课程设计的开发者提供了从数据到模型的完整闭环。压缩包内含2000个文件主体为1994张交通标志PNG图片可用于模型训练与验证另附3个Python脚本包含训练与单张图片测试代码3个XML文件作为配置或标注补充整体约201.95MB帮助理解卷积神经网络在图像分类中的典型流程。通过命令行指定训练集、测试集与模型保存路径即可快速启动训练也能用预测脚本对单张图片进行推理并查看分类结果便于调试、演示或二次开发。已有78人学习适合作为图像识别项目从零复现的参考范例。资源目录结构清晰读者可在此基础上调整网络结构、扩充数据类别进一步贴近实际智慧交通业务的落地需求也为课程答辩和项目展示提供完整素材。1. 交通标志分类为什么值得用 CNN 复现先说一个反直觉的结论在交通标志分类这个任务上模型大小不是瓶颈数据组织和预处理对齐才是。GTSRB德国交通标志基准只有 43 类每张图缩到 32×32 之后一个两层卷积的小网络就能跑到 95% 以上准确率反而是训练脚本和预测脚本的预处理不一致会让同一个模型掉到 70% 以下。这套项目实践把目录即标签的数据集、CNN 训练脚本和单图推理脚本串成了完整链路。训练时用 train.py 跑出 traffic_sign.model再用 predict.py 对任意一张标志图输出类别和置信度正好覆盖人工智能大作业和智慧交通视觉基线最常见的两种需求。适合三类人做课程设计或毕业设计的学生、智能车竞赛视觉组的备赛队伍以及第一次接触图像分类、想验证自己对卷积网络理解是否到位的工程师。下面从数据集目录讲起按训练、推理、排错的顺序把整条链路拆开。2. GTSRB 目录约定与 train.py 数据加载流程2.1 目录即标签0000000042 的编码规则GTSRB 是交通标志分类实验的默认起点43 个类别覆盖限速、禁令、警告、指示四大类。它的标签不依赖 CSV而是直接体现在目录名上train/ 和 test/ 下面各放 00000 到 00042 共 43 个子目录目录名就是类别 ID加载时把目录名转成 int 就是标签。压缩包里能看到 01146_00000.png、01639_00002.png 这类文件名前一段是图片在原始采集阶段的内部编号后一段是同类样本序号。注意 test/00000/ 目录下的文件名前缀是 00017和目录名 00000 对不上说明这个包经过重新归类所以解析标签永远以目录名为准解析文件名是常见误用换一个数据版本就会错。这个数据集还有个典型坑test/00000/ 下存在 00017_00000.png.png 这种双重扩展名文件。用 fname.endswith(.png) 过滤的话这类文件会被直接漏掉用 os.path.splitext 取扩展名则天然兼容因为 splitext 永远切最后一个点import os import cv2 import numpy as np def load_dataset(data_dir, target_size(32, 32)): images, labels [], [] for label_dir in sorted(os.listdir(data_dir)): label_path os.path.join(data_dir, label_dir) if not os.path.isdir(label_path): continue for fname in os.listdir(label_path): if os.path.splitext(fname)[1] ! .png: continue img cv2.imread(os.path.join(label_path, fname)) if img is None: continue img cv2.resize(img, target_size) images.append(img) labels.append(int(label_dir)) return np.array(images), np.array(labels)这段代码是两层循环外层 sorted 遍历类别目录保证 00000、00001 按字典序读入images 和 labels 的对应关系不乱内层按扩展名过滤splitext 拿到的总是最后一个扩展名所以 .png.png 也能通过。cv2.imread 读出来是 HWC 排布的 BGR 图resize 到统一尺寸后进数组某个文件损坏时 imread 返回 None直接跳过而不是让整个训练崩溃。归一化我一般用 img.astype(np.float32) / 255.0把像素从 0255 压到 01。也有人用 (x - mean) / std 做标准化小网络上两者差别不大但必须全项目统一训练做了哪套predict.py 就要做哪套这是第 4 章要展开的关键点。2.2 train.py 两个数据参数和一个保存参数train.py 的命令行里同时出现 --data_train 和 --data_test说明脚本设计是每个 epoch 结束就用 test 目录跑一次验证既当验证集又当最终测试集。GTSRB 官方已经划分好数据集不需要再手动 split这种设计在 43 类、每类几十到几百张的小数据集上够用。参数取值作用--data_train./train训练图像根目录含 0000000042 子目录--data_test./test每轮验证图像根目录不参与梯度更新--modeltraffic_sign.model训练完成后模型保存路径python train.py --data_train ./train --data_test ./test --model traffic_sign.model三个参数的职责边界很清楚--data_train 指定训练图像来源--data_test 指定验证图像来源--model 决定模型文件写到哪。如果 --data_test 目录不存在脚本应该在训练前报错而不是静默跳过验证否则打印出来的准确率只是模型对训练集的记忆程度。提示换成自己的数据集时验证目录要单独保留只用于算准确率绝不参与梯度更新。模型文件的保存方式有一个兼容性细节torch.save 可以存整个模型对象也可以只存 state_dict。项目里更常见的是后者加载时先实例化模型再 load_state_dict换机器、换 PyTorch 版本都不容易反序列化失败。保存时把输入尺寸和类别数一并写进 checkpointpredict.py 加载时就能对照校验省掉跨脚本排错的时间。3. CNN 卷积层配置与 traffic_sign.model 训练参数3.1 小卷积核叠加比大卷积核划算交通标志是先验极强的图像图形是标准矢量风格颜色块边界清晰干扰主要来自光照、遮挡和运动模糊。这种数据不需要 ResNet 那么深的残差结构VGG 风格的小卷积核堆叠就足够。两个串联的 3×3 卷积等价于一个 5×5 卷积的感受野参数量却少约 28%中间还多一次 ReLU 非线性对边缘和颜色边界的表达能力更强。一个典型的 TrafficNet 结构32×32 RGB 输入如下阶段层配置输出尺寸说明特征提取Conv2d(3, 32, 3, padding1) ReLU32×32×32提取边缘、颜色块特征提取Conv2d(32, 32, 3, padding1) ReLU32×32×32特征重组下采样MaxPool2d(2)16×16×32降分辨率扩大感受野特征提取Conv2d(32, 64, 3, padding1) ReLU16×16×64升通道特征提取Conv2d(64, 64, 3, padding1) ReLU16×16×64深层特征下采样MaxPool2d(2)8×8×64压缩到分类尺寸分类Linear(4096, 128) ReLU Dropout128特征到类别映射分类Linear(128, 43)43输出 logits两次 MaxPool 之后特征图是 8×8×64展平 4096 维接 128 维全连接再输出 43。这个容量对 GTSRB 恰好再深容易过拟合再浅则区分不了 33右转、34左转这类只有箭头方向不同的标志。3.2 train.py 训练循环与超参对应的 PyTorch 实现import torch import torch.nn as nn class TrafficNet(nn.Module): def __init__(self, num_classes43): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(32, 32, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(64, 64, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 8 * 8, 128), nn.ReLU(inplaceTrue), nn.Dropout(0.5), nn.Linear(128, num_classes), ) def forward(self, x): return self.classifier(self.features(x))nn.Sequential 把卷积、激活、池化按序排好forward 里先过 features 再进 classifier。Conv2d 第一个参数 3 是输入通道数对应 RGB如果数据加载改成灰度图这里必须同步改成 1否则第一个卷积直接报维度错误。Dropout 放在分类层而不是特征层因为 4096 维全连接是过拟合风险最高的地方。超参按数据集规模和脚本接口推断常见一组是 batch_size 64、Adam lr1e-3、epoch 20。GTSRB 训练集约 3.9 万张一个 epoch 大约 600 个 batchCPU 上几分钟跑完GPU 上几十秒。train.py 每轮验证的代码大致如下criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(20): model.train() for imgs, labels in train_loader: optimizer.zero_grad() preds model(imgs) loss criterion(preds, labels) loss.backward() optimizer.step() val_acc evaluate(model, test_loader) print(fepoch {epoch:02d} val_acc {val_acc:.3f}) torch.save({ state_dict: model.state_dict(), num_classes: 43, input_size: (32, 32), }, traffic_sign.model)CrossEntropyLoss 内部已经做了 softmax所以模型最后一层输出裸 logits 即可不要在 forward 里提前 softmax否则数值范围和梯度回传都会出问题。保存时把 state_dict、类别数、输入尺寸打包成字典比裸存 state_dict 多不了几个字节但 predict.py 加载时能直接校验配置一致性这是工程上值得保留的习惯。3.3 类别不均衡先跑基线再看要不要加权GTSRB 各类样本数并不均匀限速类样本明显多于部分禁令类直接训练时模型会偏向多数类。最省事的做法是在 DataLoader 里加 WeightedRandomSampler权重按类别样本数的倒数计算或者把 class_weights 传给 CrossEntropyLoss效果接近。提示类别不均衡对整体准确率影响往往不大因为多数类贡献大但少数类在混淆矩阵里被吞的情况很常见。先跑一版不加权的基线再看混淆矩阵决定是否加权别一上来就加否则无法判断提升来自权重还是来自其他改动。4. predict.py 单图推理与预处理对齐4.1 命令行参数与推理主流程推理命令是python predict.py --model traffic_sign.model -i ./test/00000/00017_00000.png.png -s--model 指向训练保存的模型文件-i 是待识别图片路径-s 表示把预测结果显示出来。这个 -s 依赖图形环境无显示的服务器上 cv2.imshow 会报错建议脚本里把 -s 实现成 cv2.imwrite 输出标注图或者用 matplotlib 保存成文件结果一样还能批量跑。对应的参数解析和推理主流程import argparse import cv2 import torch import numpy as np parser argparse.ArgumentParser() parser.add_argument(--model, requiredTrue) parser.add_argument(-i, --image, requiredTrue) parser.add_argument(-s, --show, actionstore_true) args parser.parse_args() def preprocess(img_path, target_size(32, 32)): img cv2.imread(img_path) img cv2.resize(img, target_size) img img.astype(np.float32) / 255.0 img torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0) return img def predict(model_path, img_path): ckpt torch.load(model_path, map_locationcpu) model TrafficNet(num_classesckpt[num_classes]) model.load_state_dict(ckpt[state_dict]) model.eval() with torch.no_grad(): x preprocess(img_path) logits model(x) probs torch.softmax(logits, dim1) return probs几个容易踩的细节torch.load 必须带 map_locationcpu否则在无 GPU 的机器上加载时会找 CUDA 设备直接抛错model.eval() 必须调用不调用的话 Dropout 仍处于训练模式同一次推理每次结果都会抖动with torch.no_grad() 告诉框架不建计算图单图推理影响不大但循环跑整个测试目录时差距明显。preprocess 里 permute(2, 0, 1) 把 HWC 变成 CHWunsqueeze(0) 加 batch 维这两步漏一个都会在 forward 时报维度错误。输入尺寸和归一化必须与第 2 章 load_dataset 完全一致这就是训练脚本和推理脚本最容易分叉的地方。4.2 预处理不一致是准确率掉落的头号原因推理结果比训练时低一截八成是预处理没对齐。常见三类不一致项训练端推理端后果resize 尺寸32×3264×64特征尺度错位准确率明显下降归一化方式/255.0(x-mean)/std数值分布不同模型输入偏移通道顺序BGRcv2RGBPIL颜色语义错乱红蓝互换通道顺序问题最隐蔽也最致命卷积网络学到的颜色权重是按训练时通道顺序排列的红蓝对调后限速标志的红色会被当成蓝色特征处理准确率大幅下跌且很难从日志里看出来。验证方法很直接挑训练集里准确率最高的一张图分别用两套预处理各跑一次推理结果不一致就说明链路有偏差。更稳妥的做法是 predict.py 直接复用 load_dataset 里的同一个预处理函数从函数层面杜绝分叉而不是在推理脚本里复制粘贴一份。4.3 topk 输出怎么映射回标志含义softmax 之后取概率前五topv, topi probs.topk(5) for score, cls_id in zip(topv[0].tolist(), topi[0].tolist()): print(fclass {cls_id:02d} prob {score:.3f})topi 的元素就是类别 ID与训练时数据目录的排序一一对应0 对应限速 201 对应限速 302 对应限速 50依此类推。predict.py 最好维护一份 {class_id: 中文标志名} 的映射表打印结果时把名称带出来否则人工核对 43 个类别 ID 很容易看错尤其在 11、12、13 这类外观接近的标志上。5. 易混淆标志的验证手段与误判排查5.1 用 classification_report 定位问题类别43 类整体准确率看不出问题要落地到智慧交通场景必须知道具体哪些标志互相认错。把 test 目录全部过一遍累积真值和预测值用 sklearn 输出分类报告和混淆矩阵from sklearn.metrics import confusion_matrix, classification_report print(classification_report(y_true, y_pred, digits3)) cm confusion_matrix(y_true, y_pred)重点看每个类别的 recall 而不是整体 accuracy。限速类 recall 低说明数字被误读禁令类 precision 低说明其他类被误判成它。把对角线上数值低的类别挑出来再去对应目录看原始图片比盯着 loss 曲线有效得多。5.2 高频误判对的典型特征GTSRB 里最容易混淆的是成对出现的标志类别对视觉相似点可靠区分线索1 限速30 / 2 限速50红圈白底数字易糊数字笔画宽度需要更高输入分辨率11 路口先行 / 12 优先道路同为黄色菱形系内部图形结构不同观察特征图激活区域33 右转 / 34 左转白色箭头方向相反箭头朝向关注中心 ROI38 靠右行驶 / 39 靠左行驶方向相反箭头位置与朝向对这类成对类别最后一个实用的手段是在预测脚本里加置信度门槛当 top1 和 top2 的概率差小于阈值比如 0.05时同时输出两个候选项并提示人工复核。真实交通场景下这比强行二选一更安全自动驾驶的决策模块也更愿意收到带不确定性的结果。5.3 不改模型也能提点的预处理技巧如果验证精度卡在 94% 左右上不去先检查数据而不是模型。把误判图片按文件名批量拷出来看常见情况是亮度极低的夜间图、被杆子遮挡一半的标志、以及缩放后数字变糊的限速牌。对应的处理是训练端加随机亮度抖动和随机裁剪推理端保持与训练一致不做额外增强。数据包里那些双重扩展名的文件顺手用一条命令清理干净能避免后续所有脚本在文件过滤上踩同样的坑find ./test -name *.png.png -exec bash -c mv $0 ${0%.png} {} \;这条命令把 .png.png 的尾缀去掉恢复成标准单扩展名之后任何按 .png 过滤的代码都能正常工作不再依赖 splitext 的特殊处理。本文还有配套的精品资源点击获取