简介这份资源面向计算机、人工智能、电子信息等相关专业的学生及企业员工提供一套基于卷积神经网络识别交通标志的完整Python项目源码适合作为课程设计、毕业设计、大作业或初期项目立项演示的参考案例也便于初学者进行实战练习。压缩包共包含9个文件以5个py源码文件为核心辅以2个csv数据索引文件、1个xml配置文件和1个md项目说明文档整体约310KB结构紧凑、便于快速上手。项目围绕GTSRB交通标志数据集展开涵盖数据预处理、模型构建、训练与评估等关键环节读者可借此理解CNN在图像分类任务中的完整实现流程掌握数据加载、网络搭建、模型训练与性能评估的排错思路并可直接运行验证效果。目前已有170人学习下载具备较高的学习借鉴价值。1. 从一张模糊的路牌说起CNN 识别交通标志到底难在哪限速 40 的圆牌被树影切掉一半逆光下红色边缘发灰摄像头还有运动模糊——这种图丢给传统模板匹配基本当场翻车。基于 CNN 识别交通标志python 源码项目说明数据集是 GTSRB这套东西解决的正是这类真实路况下的分类问题输入一张裁剪好的标志小图输出它属于 43 个类别中的哪一个。GTSRB 全称 German Traffic Sign Recognition Benchmark是交通标志识别领域被引用最多的公开数据集之一六万多张实拍图覆盖限速、禁令、指示、警告等类别类别极不均衡小类只有两百来张大类上千张。这套方案适合三类人刚学完 cnn 卷积神经网络想找个完整项目练手的入门者需要快速搭一个交通标志分类基线再迭代的算法工程师以及做嵌入式视觉、想把模型压到能上车的开发者。它不解决检测只解决分类这个边界先划清楚。2. GTSRB 数据集的真实结构与预处理链路2.1 为什么不能直接 ImageFolder 一把梭很多人拿到 GTSRB 压缩包第一反应是torchvision.datasets.ImageFolder直接读然后发现读不出来或者类别全乱。原因是 GTSRB 的目录组织方式和标准 ImageFolder 不一样。它通常长这样训练集是Train/00000/到Train/00042/共 43 个文件夹每个文件夹里是00000_00000.ppm这种命名的图片外加一个GT-00000.csv记录每张图的 ROI 坐标。测试集是Final_Test/Images/一个大目录配一个GT-final_test.csv记录文件名和标签。也就是说测试集没有按类别分文件夹ImageFolder 对它无效。常见做法是写一个自定义 Dataset把 CSV 读进来按 ROI 裁剪再统一 resize。ROI 这一步别省GTSRB 的图四周有大量无关背景直接整图训练会让模型学到背景捷径测试集一换场景就崩。import os import pandas as pd from PIL import Image from torch.utils.data import Dataset from torchvision import transforms class GTSRBDataset(Dataset): def __init__(self, root, csv_file, transformNone, use_roiTrue): # root: 图片所在目录; csv_file: 标签文件路径 self.root root self.df pd.read_csv(csv_file, sep;) # GTSRB 的 CSV 用分号分隔 self.transform transform self.use_roi use_roi def __len__(self): return len(self.df) def __getitem__(self, idx): row self.df.iloc[idx] img_path os.path.join(self.root, row[Filename]) img Image.open(img_path).convert(RGB) if self.use_roi: # ROI 四角坐标裁剪掉背景 x1, y1, x2, y2 row[Roi.X1], row[Roi.Y1], row[Roi.X2], row[Roi.Y2] img img.crop((x1, y1, x2, y2)) label int(row[ClassId]) if self.transform: img self.transform(img) return img, label逻辑说明sep;是关键GTSRB 的 CSV 默认分号分隔用逗号会读成单列。use_roi控制是否裁剪训练时建议开推理部署时如果上游已经给了裁剪好的图就关掉。参数上Roi.X1这些列名在不同版本 CSV 里大小写可能不同读之前先print(self.df.columns)确认一遍这是血泪经验。2.2 归一化参数与数据增强的取舍预处理里最容易拍脑袋的是均值和方差。GTSRB 图片偏暗、对比度低直接用 ImageNet 的mean[0.485,0.456,0.406]不是不行但自己统计一遍更稳。统计方法很简单遍历训练集算每个通道的均值和标准差写死到 transform 里。# 统计训练集通道均值方差跑一次即可结果写死 from torch.utils.data import DataLoader stats_transform transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), ]) ds GTSRBDataset(Train, Train.csv, transformstats_transform) loader DataLoader(ds, batch_size256, shuffleFalse) mean 0.0; std 0.0; n 0 for imgs, _ in loader: b imgs.size(0) imgs imgs.view(b, 3, -1) mean imgs.mean(2).sum(0) std imgs.std(2).sum(0) n b print(mean / n, std / n)数据增强方面交通标志有个特殊性水平翻转要慎用。左转和右转标志翻转后语义就反了限速数字翻转后也不成立。我一般只开轻度旋转±10 度、亮度对比度抖动、小范围平移缩放翻转留给那些本身对称的类别或者干脆不开。这个坑很多人踩过训练集准确率很高测试集一塌糊涂查半天发现是翻转把标签语义搞乱了。3. 用 PyTorch 搭一个能跑通的 CNN 分类器3.1 网络结构从 LeNet 改起还是直接上 ResNet标题里说的是 CNN没指定具体架构。我的建议是分两步先用一个自己搭的小 CNN 把流程跑通确认数据管道没问题再换 ResNet 之类的预训练模型冲精度。小 CNN 参考 LeNet 改输入 32x32两个卷积块加全连接参数量几十万CPU 上都能训。import torch import torch.nn as nn class TrafficCNN(nn.Module): def __init__(self, num_classes43): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.Conv2d(32, 32, 3, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(2), # 32 - 16 nn.Conv2d(32, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 64, 3, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2), # 16 - 8 nn.Conv2d(64, 128, 3, padding1), nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2), # 8 - 4 ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(128 * 4 * 4, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes), ) def forward(self, x): return self.classifier(self.features(x))逻辑说明每个卷积后接 BatchNorm 和 ReLUBN 放在卷积和激活之间是标准顺序能明显加快收敛。三次池化把 32x32 降到 4x4最后展平接全连接。Dropout(0.5)放在全连接前防止小数据集过拟合。参数上num_classes43是 GTSRB 的类别数改数据集时记得同步改。如果显存够把通道数翻倍精度会更好但训练时间也翻倍。3.2 训练循环与学习率调度训练循环本身不复杂关键是几个参数优化器用 Adam初始学习率 1e-3配合余弦退火或者 StepLR。损失函数用 CrossEntropyLossGTSRB 类别不均衡可以加weight参数给样本少的类更高权重。from torch.optim import Adam from torch.optim.lr_scheduler import CosineAnnealingLR from torch.utils.data import DataLoader from torchvision import transforms train_tf transforms.Compose([ transforms.Resize((32, 32)), transforms.RandomRotation(10), transforms.ColorJitter(brightness0.3, contrast0.3), transforms.ToTensor(), transforms.Normalize(mean[0.34, 0.31, 0.32], std[0.27, 0.26, 0.27]), # 换成自己统计的值 ]) train_ds GTSRBDataset(Train, Train.csv, transformtrain_tf) train_loader DataLoader(train_ds, batch_size128, shuffleTrue, num_workers4) device torch.device(cuda if torch.cuda.is_available() else cpu) model TrafficCNN(num_classes43).to(device) criterion nn.CrossEntropyLoss() optimizer Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max30) for epoch in range(30): model.train() total_loss 0.0 for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() out model(imgs) loss criterion(out, labels) loss.backward() optimizer.step() total_loss loss.item() scheduler.step() print(fepoch {epoch}, loss {total_loss / len(train_loader):.4f})逻辑说明weight_decay1e-4是 L2 正则配合 Dropout 一起压过拟合。CosineAnnealingLR的T_max设成总 epoch 数学习率从 1e-3 余弦降到接近 0。num_workers4在 Linux 上没问题Windows 上如果报错就改成 0。训练时盯着 loss 曲线如果前几个 epoch loss 不降先查数据管道八成是标签对不上或者归一化参数写错了。3.3 在测试集上评估并生成混淆矩阵GTSRB 的测试集评估要读GT-final_test.csv按文件名匹配标签。评估时记得model.eval()和torch.no_grad()否则 BN 层和 Dropout 行为不对结果会偏低。import numpy as np from sklearn.metrics import confusion_matrix, classification_report test_tf transforms.Compose([ transforms.Resize((32, 32)), transforms.ToTensor(), transforms.Normalize(mean[0.34, 0.31, 0.32], std[0.27, 0.26, 0.27]), ]) test_ds GTSRBDataset(Final_Test/Images, GT-final_test.csv, transformtest_tf) test_loader DataLoader(test_ds, batch_size256, shuffleFalse) model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in test_loader: imgs imgs.to(device) preds model(imgs).argmax(1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, digits4)) cm confusion_matrix(all_labels, all_preds)逻辑说明argmax(1)取每行最大值的索引作为预测类别。classification_report会输出每个类的 precision、recall、f1重点看样本少的类 recall 是不是特别低。混淆矩阵里如果发现某两类互相混淆严重比如限速 30 和限速 50说明模型对数字细节分辨不够可以考虑提高输入分辨率或者加注意力模块。参数上测试集 batch_size 可以开大因为不需要反向传播显存占用小。4. 训练过程中的避坑与排查清单4.1 准确率卡在 60% 上不去现象训练 loss 正常下降但验证准确率到 60% 左右就横盘怎么调学习率都没用。原因通常是输入分辨率太低。GTSRB 里很多标志的区分点在于数字和细小图案32x32 下限速 30 和 50 几乎糊成一团。解决把输入提到 48x48 或 64x64网络结构相应调整准确率通常能涨 5 到 10 个百分点。代价是训练变慢自己权衡。4.2 测试集准确率远低于训练集现象训练集 99%测试集 85%差距十几个点。原因有两个常见来源一是数据增强里的水平翻转把语义搞反了二是归一化参数用了 ImageNet 的而不是自己统计的。解决先关掉翻转重训一遍看差距是否缩小再把归一化换成训练集统计值。如果还不行检查测试集的 ROI 裁剪是否和训练集一致不一致会导致输入分布偏移。4.3 某些类别 recall 极低现象整体准确率 95%但某几个类 recall 只有 50%。原因类别不均衡小类样本太少模型倾向于预测大类。解决在 CrossEntropyLoss 里加weight权重按类别频率的倒数算或者用重采样让每个 batch 里小类占比提高。注意权重别设太极端否则大类性能会掉。4.4 训练中途 loss 变成 nan现象跑着跑着 loss 突然 nan梯度爆炸。原因学习率太大或者数据里有脏图全黑、损坏。解决先把学习率降到 1e-4 试试如果还 nan写个脚本遍历数据集检查有没有异常图比如像素全 0 或者尺寸为 0 的。另外 Adam 的 eps 默认 1e-8一般够用但如果输入数值范围很大可以适当调大。4.5 推理时单张图预测结果不稳定现象同一张图单独推理和批量推理结果不一样。原因忘了model.eval()BN 层在训练模式下用 batch 统计量单张图时统计量不准。解决推理前一定加model.eval()并用torch.no_grad()包住。这个坑很隐蔽因为批量推理时 batch 够大BN 统计量接近全局值问题不明显单张就暴露了。5. 把模型推到 99% 以上几个我常用的进阶技巧第一个技巧是迁移学习。自己搭的小 CNN 到 97% 左右就顶天了想再往上直接上 ResNet18 或 EfficientNet-B0加载 ImageNet 预训练权重把第一层改成适应 GTSRB 的输入最后一层改成 43 类。训练时前几层冻结只训后面学习率设小一点1e-4 起步。我实测 ResNet18 在 64x64 输入下能到 99.2%比小 CNN 高两个点训练时间也就多一倍。import torchvision.models as models model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 43) # 替换分类头 # 冻结前面的层 for name, param in model.named_parameters(): if fc not in name: param.requires_grad False model model.to(device)第二个技巧是测试时增强TTA。推理时对同一张图做几种轻微变换比如不同亮度、小角度旋转把预测概率平均能稳定涨 0.3 到 0.5 个点。代价是推理时间翻几倍看场景取舍。方案输入尺寸测试集准确率单张推理耗时GPU自建小 CNN32x3296.8%2ms自建小 CNN64x6498.1%5msResNet18 微调64x6499.2%8msResNet18 TTA64x6499.5%30ms第三个技巧是错误分析驱动迭代。别盲目调参先把混淆矩阵里错得最多的类对找出来看看那些图长什么样。我上次发现限速 80 和限速 80 解除混淆严重原因是两者图案几乎一样只差一条斜杠模型在低分辨率下分不出来。把这两类的图单独拿出来提高分辨率重训问题就解决了。这个思路比无脑加数据增强有效得多。最后一个习惯每次实验都记配置。学习率、batch size、增强方式、输入尺寸、随机种子全写进一个 csv 或者文本里。我吃过亏某次跑出一个好结果回头想复现忘了当时用的什么增强参数重跑三遍都对不上。现在我用一个简单的实验记录表每跑一次填一行省了太多后悔药。这套基于 CNN 识别交通标志的方案从数据管道到模型训练再到调优链路是完整的值得动手跑一遍。希望帮到你。本文还有配套的精品资源点击获取
