恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

基于CNN的交通标志识别:GTSRB数据集与TSR-master项目实战

  • 首页
  • 资讯中心
  • /
  • 基于CNN的交通标志识别:GTSRB数据集与TSR-master项目实战

相关资讯

FedAvg联邦学习实战:用MNIST手写数字识别跑通完整流程 2026/9/28 16:27:47
轮胎DOT编码识别:工业OCR鲁棒性实战指南 2026/9/28 16:27:47
STM32音乐播放器实战:WAV解析与PWM/DAC音频输出 2026/9/28 16:27:47

最新资讯

ax:用一条CLI命令将agentic工作负载调度到Kubernetes
ax技术定位解析:从Kubernetes调度到Agent Substrate实践
ax:用CLI将agentic工作负载调度到Kubernetes的实践指南
TMS320F280049C ePWM死区控制详解:寄存器配置与完整代码实现
Hono + pm2 部署后接口无响应?从端口监听到进程守护的排查指南
Python网络舆情分析系统实战:从环境搭建到情感分析可视化全流程

今日推荐

婚恋网站实战案例:避开3个高价坑,省钱50%还能跑赢流量
制作网页比较方便的软件怎么选?一文搞懂避坑指南
BootCamp6.1.7071驱动包手动安装与回滚全攻略

本周热门

从像素到笔画:srt-whiteboard-animation骨架笔迹追踪实现(Zhang-Suen细化+8邻接追踪)
网站建设的英语怎么说?别只背单词,看完这套安全完整流程才敢上线
新手入门看这篇:建设网站加盟避坑指南与SEO实操

本月精选

自研推理加速器Redwood:两周内实现PyTorch模型高效部署的实战教程
V4L2摄像头采集实战:从camera_client.rar到出图全流程解析
从“谁发明了钢琴键”到知识问答智能体:RAG与记忆工程实践

基于CNN的交通标志识别:GTSRB数据集与TSR-master项目实战

发布时间:2026/9/28 16:27:47
基于CNN的交通标志识别:GTSRB数据集与TSR-master项目实战 简介这是一份面向智慧交通场景的CNN交通标志识别实践项目核心借助GTSRB数据集完成从数据预处理、模型构建到训练评估的全流程适合有一定Python与深度学习基础的学习者作为课程设计或项目练手。资源压缩包约310KB共8个文件包括5个Python脚本、2个CSV文件与1个XML配置文件脚本覆盖数据输入、预处理、CNN网络定义、训练与评估等完整环节CSV文件分别存放训练集与测试集的样本索引XML用于项目配置。GTSRB数据集包含德国道路上43类交通标志、约5万张标注图片可有效检验模型泛化能力。项目代码按标准深度学习流程组织结构清晰可直接运行并替换数据集进行二次开发。目前已有623人学习下载适合希望快速掌握CNN图像分类实战流程并落地交通标志识别场景的开发者参考。1. 用CNN识别交通标志GTSRB数据集与TSR-master项目完整拆解交通标志识别是智慧交通里最不好糊弄的图像任务之一同一个限速牌白天逆光、晚上反光、雨天溅泥、运动模糊形态差异远超想象。传统CV规则写不到头HOG加SVM也吃不住这种噪声。这个项目走CNN路线对GTSRB数据集的43类德国交通标志做分类预处理、网络定义、训练评估一条线跑通。压缩包TSR-master目录下是Preprocessing.py、TSRInput.py、TSRCnn.py、TSRTrain.py、TSREval.py五个核心脚本外加data3图像目录和train_data.csv、test_data.csv标注文件。数据和代码拆分得干净每个环节都能单独改。适合正在做人工智能大作业、毕设或想拿真实数据集练手CNN图像分类的从业者。代码不是那种能看不能跑的演示改改数据路径就能复现。2. 先读懂GTSRB和项目结构43类数据的分隔符、切分方式与文件职责2.1 GTSRB数据集43类标注怎么组织切分时最容易翻车的点GTSRBGerman Traffic Sign Recognition Benchmark是德国交通标志识别基准43个类别覆盖限速、禁止通行、转弯、施工这些常见标志训练与测试图像合计约5万张分辨率不高但干扰丰富。项目里的data3目录放的就是这些图像train_data.csv和test_data.csv记录标注。csv的格式是每行一条记录包含图像文件名和类别编号类别编号从0到42。这里有一个必须注意的点原始GTSRB的csv用的是分号分隔而不是逗号。如果直接用pandas默认的sep,去读整行会被当成一列后面按列取值索引全部错位训练脚本一启动就报KeyError。我一般会用sep;显式声明保险起见加encodingutf-8Windows环境下偶尔会有编码坑。读出来之后先确认列名常见的是Filename和ClassId两列如果命名不完全一致就按实际列名改代码里的字符串。import pandas as pd # 读入训练和测试标注GTSRB官方csv是分号分隔 train_df pd.read_csv(data3/train_data.csv, sep;, encodingutf-8) test_df pd.read_csv(data3/test_data.csv, sep;, encodingutf-8) print(train_df.head()) # 确认列名和路径格式 print(train_df[ClassId].nunique()) # 期望输出43代表全类别逻辑说明pandas按sep参数拆列所以分隔符错了后面全错nunique()用来核对类别总数防止csv解析错或数据不全。如果你手里的csv列名是class_id而不是ClassId把后面的取值改一下即可。train_df和test_df分别对应训练和测试标注这两个DataFrame在TSRInput.py里还要重复用到所以最好单独跑一遍确认能正常读。另一个容易翻车的是训练/验证切分。原始GTSRB数据里同一个交通标志的连拍帧是连续存放的如果直接对整个目录随机shuffle再切训练验证同一物理标志的相邻帧会同时出现在train和val里验证集准确率虚高。常见做法是按图像所属的原始场景先分组再把整组切到一边去。这种细节直接影响评估可信度训练代码再漂亮切分错了也是白搭。这里还要注意类别分布。GTSRB里限速类标志样本上千少数警告标志可能只有一两百张如果不处理模型会偏向样本多的类。常见做法是给DataLoader配加权采样器或者训练时用带class_weight的损失函数让少数类的梯度贡献更大。2.2 为什么是CNN而不是HOGSVM或MLP交通标志识别的难点不在标志长什么样而在噪声下还能不能认出它。GTSRB采集自真实道路光照方向、曝光差异、遮挡、运动模糊、残影都有同一个类别内差异非常大。HOG特征加线性SVM是传统方案里的经典组合对固定姿态和清晰边缘有效但视角和光照变化一旦超出特征设计时假设的范围准确率掉得厉害。手工特征本质上是在用人的经验猜哪些模式重要在GTSRB这种数据里猜不准。CNN把特征提取交给网络自己。卷积层通过局部连接和权值共享扫描整张图低层学到边缘、角点、颜色块高层把这些组合成形状语义路径牌、限速数字、红圈斜杠都能逐层抽象出来。池化带来平移容忍性同一个标志稍微移位、缩放网络依然能识别。对GTSRB这种类别多、噪声宽的数据CNN的结构性优势是明确的。从训练成本看GTSRB是CNN入门的友好数据集。图像尺寸小单张GPU卡上小网络一个epoch几分钟到十几分钟个人电脑跑得动不像ImageNet级别训练动辄以天为单位。这也是很多人工智能大作业和课程设计选它的原因。有一点值得说明这个项目的TSRCnn.py属于浅层卷积网络不是ResNet那种几十层结构。对交通标志识别浅层网络已经能提取足够判别性的特征且训练快、显存占用低、不容易过拟合小数据集。想换ResNet18替换模型定义即可训练管线不用动。很多人会问为什么不用全连接网络直接分类。MLP把每个像素当成独立特征二维邻域结构全部丢失同样的标志平移几个像素就认不出来。CNN的卷积操作天然保留空间结构这是两者在图像任务上差距悬殊的根本原因。SVM加手工特征做得再精细也逃不过特征工程的天花板——你只能提取你想到的特征而CNN能自己找到你没想过的特征。2.3 TSR-master项目文件架构五个脚本各自管什么项目主目录TSR-master的核心是五个python文件加上data3图像目录和两个csv标注。从命名看职责划分很干净TSRInput负责数据输入TSRCnn负责模型定义TSRTrain负责训练TSREval负责评估Preprocessing负责预处理。文件职责关键输入/输出Preprocessing.py图像缩放、归一化、数据增强输入原始图像目录输出处理后的numpy数组或transform对象TSRInput.py定义Dataset类读取csv并加载图像输入csv路径图像目录输出(张量,标签)对TSRCnn.py定义CNN网络结构输入图像张量输出43维分类logitsTSRTrain.py训练循环、损失计算、模型保存输入训练/验证DataLoader输出训练好的模型文件TSREval.py在测试集上评估模型输入权重测试数据输出准确率、混淆矩阵这种按环节拆文件的组织方式在项目实践里很常见好处是每个环节能单独改。数据目录和csv路径通常写死在脚本里下载后第一步打开TSRTrain.py看它引用的路径和实际目录是否一致不一致改一处就行。运行顺序是数据放好 → 跑Preprocessing做增强 → TSRTrain训练并保存权重 → TSREval评估。.idea和vcs.xml是PyCharm和版本控制配置与模型逻辑无关。调试的时候我习惯先跑一个最小路径不加载全部数据只取200张图、训练1个epoch确认Loss在下降、DataLoader不报错再全量跑。这一步能省下大量排错时间尤其是换机器之后路径、依赖库版本全对不上的时候。3. 数据预处理与输入管道从csv路径到模型能吃到的张量3.1 Preprocessing.py 的预处理逻辑尺寸、归一化和增强CNN不能直接吃任意尺寸的原始图片。GTSRB原图尺寸不统一有的只有几十像素有的上百像素网络输入层固定尺寸第一步就是把所有图缩放到同一个大小。这个项目里比较稳的选择是32x32或48x48。32x32速度快48x48保留更多细节对限速牌上的数字识别更友好。我自己习惯用48x48因为GTSRB很多标志细节靠数字和字符区分缩到32有些细纹理就糊了。预处理一般包含三步转成RGB三通道数组原图如果是灰度就复制三遍、缩放到统一尺寸、归一化到[0,1]或按均值和标准差标准化。PyTorch的常见写法是借助torchvision.transforms一步到位from torchvision import transforms # 训练用预处理缩放随机旋转归一化 train_transform transforms.Compose([ transforms.Resize((48, 48)), transforms.RandomRotation(degrees10), transforms.ToTensor(), transforms.Normalize(mean[0.340, 0.312, 0.321], std[0.202, 0.199, 0.201]) ]) # 验证和测试用只做尺寸调整和归一化不做随机增强 eval_transform transforms.Compose([ transforms.Resize((48, 48)), transforms.ToTensor(), transforms.Normalize(mean[0.340, 0.312, 0.321], std[0.202, 0.199, 0.201]) ])逻辑说明Resize统一尺寸RandomRotation是数据增强。水平翻转在交通标志场景要慎用——如果标志本身带方向性比如左转箭头翻转会把语义反掉。我一般只用旋转增强或者只对对称类别的标志开翻转。Normalize里那组mean/std是从GTSRB训练集统计的近似值不想纠结就用官方常用的(0.5,0.5,0.5)也能跑。注意验证集和测试集绝对不能用随机增强。eval_transform里不加随机旋转和翻转原因很简单——验证集加增强会让每个epoch的结果不同你不知道这个epoch的acc是模型变好了还是随机增强叠出来的。所有正规实验里验证集都是确定性变换。参数说明RandomRotation的degrees10表示在-10度到10度之间随机旋转这个幅度对交通标志安全。数据增强不是越多越好过度的颜色抖动、透视扭曲会破坏标志的几何清晰度。我见过有人给GTSRB加强颜色抖动和Cutout结果训练集acc很高测试集不升反降。颜色空间方面也有讲究。GTSRB里有少量灰度图但大部分是彩色交通标志的颜色本身就是判别信息——红色禁令、蓝色指示、黄色警告。所以一般保留RGB三通道不要转灰度。如果某些原图是灰度模式在读取时统一convert(RGB)保证通道数一致。3.2 TSRInput.py数据集类与DataLoader的写法PyTorch里TSRInput.py的核心是继承torch.utils.data.Dataset的类。它的任务有两块从csv里读出文件名和标签然后在__getitem__里按索引加载图像并应用预处理。写得好的Dataset类训练时DataLoader能按batch批量取数、并行加载不卡IO。import os import pandas as pd from PIL import Image from torch.utils.data import Dataset class GTSRBDataset(Dataset): def __init__(self, csv_path, img_dir, transformNone): self.df pd.read_csv(csv_path, sep;, encodingutf-8) self.img_dir img_dir self.transform transform def __len__(self): return len(self.df) def __getitem__(self, idx): row self.df.iloc[idx] img_path os.path.join(self.img_dir, row[Filename]) image Image.open(img_path).convert(RGB) label int(row[ClassId]) if self.transform: image self.transform(image) return image, label逻辑说明__getitem__里按idx取出csv中的一行用os.path.join拼接图像路径PIL打开后转RGB。转RGB这步很关键GTSRB里有灰度图不转换的话同一批张量通道数不统一模型输入维度会炸。label从csv读出后转int因为有些csv的类别是带引号的字符串。参数说明csv_path和img_dir按你的目录结构传。如果data3下的图按43个子文件夹组织Filename存的是相对路径os.path.join恰好能拼上如果csv里存的是绝对路径就把这行改成直接用row[Filename]。内存方面GTSRB每张图几十KB全量读进内存大概几个GB8G内存的机器建议用上面的逐张读取方案不要一次性把所有图load到list。如果你机器内存够大可以在Dataset里加一个cache字典第一次读图后存起来第二次直接命中训练速度能快不少。Dataset类写好之后在TSRTrain.py里用DataLoader包一层from torch.utils.data import DataLoader train_dataset GTSRBDataset(data3/train_data.csv, data3, train_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers2, drop_lastTrue)逻辑说明shuffleTrue保证每个epoch样本顺序不同防止模型学到batch之间的顺序相关性。num_workers是读取图像的并行进程数Windows下建议设0或1多进程在某些环境下会反复报DataLoader worker错误这是PyTorch在Windows上的老问题。drop_lastTrue是因为最后的batch如果凑不够batch_sizeBatchNorm层计算会有边界问题。参数说明batch_size32是比较稳的起步值显存不够降到16训练太慢提到64。对GTSRB这种5万级别数据64以上的batch会让梯度平均过于光滑收敛反而变慢。4. 模型构建与训练TSRCnn.py的网络定义和TSRTrain.py的训练循环4.1 TSRCnn.py 的CNN结构卷积、池化、全连接怎么组合TSRCnn.py里定义的网络核心是卷积ReLU池化堆叠再接全连接层输出43个类别的得分。对GTSRB这个规模常见的结构是三层卷积加两层全连接。第一层学边缘和色块第二层学纹理组合第三层学部件结构最后全连接做分类。import torch.nn as nn class TSRCNN(nn.Module): def __init__(self, num_classes43): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), nn.Conv2d(64, 128, kernel_size3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(128 * 6 * 6, 512), nn.ReLU(inplaceTrue), nn.Dropout(p0.5), nn.Linear(512, num_classes), ) def forward(self, x): return self.classifier(self.features(x))逻辑说明输入是3x48x48的图像。三个卷积层通道数32→64→128每层后接ReLU和2x2最大池化。48的尺寸经三次池化变成6x6全连接第一层输入维度就是12866。这个维度是算出来的不是随便写的——改输入尺寸或池化次数这里必须同步改否则forward时直接shape mismatch。Dropout0.5放在全连接前防过拟合如果训练集acc远高于验证集可以把它提到0.6。参数说明kernel_size3和padding1组合保证卷积不改变特征图尺寸尺寸只由池化控制。这里没用BatchNorm浅层网络加不加差别没那么大如果换成ResNet这种深层结构BatchNorm就是必须的了。关于感受野三次3x3卷积加两次池化后最后一个卷积层每个神经元大约能看到原始图像里一块不小的区域足够覆盖交通标志的核心图形。如果只用一层卷积感受野太小模型只能看到局部纹理看不到整体形状识别率会明显下降。这也是为什么这个结构至少要三层卷积——GTSRB里的标志是靠整体形状和内部字符区分的。这个网络大概几十万个参数训练很快。但也不是越深越好GTSRB训练集就5万张网络太深容易把噪声也背下来。我试过用预训练ResNet迁移到这个任务提升非常有限训练时间翻了好几倍原因就是数据量不足以支撑深层网络的微调。4.2 TSRTrain.py 训练循环损失函数、优化器和学习率训练循环是TSRTrain.py的核心逻辑包含前向传播、损失计算、反向传播、参数更新四个固定动作。损失函数用交叉熵优化器用Adam是常见选择学习率从0.001起步。import torch import torch.nn as nn import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model TSRCNN(num_classes43).to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) num_epochs 30 for epoch in range(num_epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) avg_loss running_loss / len(train_loader.dataset) print(fEpoch {epoch1}/{num_epochs}, Loss: {avg_loss:.4f})逻辑说明optimizer.zero_grad()先清空上一步梯度PyTorch梯度默认累加loss.backward()算梯度optimizer.step()更新权重。交叉熵内部已经把softmax和log合并了网络最后一层输出裸logits即可不需要手动接softmax。model.train()切换训练模式影响Dropout和BatchNorm的行为。参数说明lr1e-3是Adam对这类小图像的稳妥起点。loss震荡不下降就降到5e-4收敛太慢就先跑5个epoch到1e-3再手动降到1e-4做精调。num_epochs30对这个网络偏保守一般15到20个epoch就收敛了后面损失不再下降就该停继续跑就是纯过拟合。关于优化器的选择SGD加动量也是这个数据集的常见配置momentum0.9、lr从0.01起步配合学习率衰减。相比之下Adam收敛快但对超参数敏感SGD最终准确率往往略高但调起来麻烦。我个人的做法是先Adam跑通流程再换SGD跑一轮对比最终精度。如果你的loss在某个epoch突然变成nan大概率是lr太大把学习率降到原来的十分之一重来。4.3 验证集与早停防止训练集acc虚高训练过程中如果没有验证集你根本不知道第几个epoch开始过拟合。常见做法是从训练数据里切出10%作为验证集每个epoch结束后在验证集上跑一次记录验证loss和准确率。TSRTrain.py保存验证集上最好的模型权重这个习惯很关键。best_acc 0.0 for epoch in range(num_epochs): # ... 训练代码同上 ... model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) val_acc correct / total if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pt)逻辑说明model.eval()关掉Dropout和BatchNorm的训练行为保证验证结果可复现。torch.no_grad()关掉梯度追踪省显存也省时间。best_acc记录历史最高验证准确率只有超过它才覆盖保存最后留下的永远是验证集上最强的那一版权重。参数说明best_model.pt是PyTorch的state_dict格式只保存参数不保存完整模型。加载时要先实例化TSRCNN再load_state_dict。如果你用的TensorFlow对应逻辑是ModelCheckpoint回调设monitorval_acc和save_best_onlyTrue效果等价。如果你想部署到生产环境建议额外保存一份完整模型torch.save(model)或者SavedModel格式都行省得推理时还要重新拼网络定义。训练时每5个epoch看一眼验证集和训练集的差值如果差值从2%慢慢拉大到10%说明模型已经开始背训练集了这时候要么调大Dropout要么提前停。5. 评估与避坑TSREval.py的指标计算和五个实测踩坑记录5.1 TSREval.py测试集准确率与混淆矩阵怎么看训练完不能只看训练loss最终评价要在测试集上做。TSREval.py加载best_model.pt在test_loader上跑完整评估。测试集和训练集完全独立这里的准确率才是模型真实泛化能力的参考值。import torch from sklearn.metrics import accuracy_score, confusion_matrix model.load_state_dict(torch.load(best_model.pt, map_locationcpu)) model.eval() all_preds, all_labels [], [] with torch.no_grad(): for images, labels in test_loader: images images.to(device) outputs model(images) _, preds torch.max(outputs, 1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.numpy()) acc accuracy_score(all_labels, all_preds) cm confusion_matrix(all_labels, all_preds) print(fTest Accuracy: {acc:.4f}) print(cm)逻辑说明map_locationcpu允许没有GPU的机器加载GPU上训练的权重。torch.max(outputs, 1)返回每行最大值和索引索引就是预测类别。confusion_matrix输出43x43矩阵对角线越亮越好哪类容易混淆一眼就能看出。参数说明如果你的模型是TensorFlow的h5格式用model.load_weights替换load_state_dict预测用model.predict其余逻辑一样。测试时DataLoader不用shuffle。光看整体准确率不够。GTSRB有些类别样本数特别少整体95%可能只是常见类全对了罕见类几乎全错。我一般会额外算每一类的召回率低的那几个类单独拉出来看混淆矩阵确认是和哪个常见类弄混了。比如禁止驶入和禁止机动车都有红圈白底就很容易互相误判。这个任务的常见表现是浅层CNN测试准确率在90%到97%之间人类水平大约99%。如果你的模型到了97%但怎么都上不去别急着堆层数先看混淆矩阵里是哪几类在互相打架。很多时候是某几个标志长得太像或者你的预处理把关键颜色信息破坏了。5.2 避坑记录五条实测踩坑与对应解法踩坑记录这块我直接按现象、原因、解决三段式写都是我实际跑GTSRB时碰过的。第一条csv读出来全部挤在一列里。现象是pandas读完后只有一列按列名取数据直接KeyError。原因是GTSRB官方csv用分号分隔pandas默认按逗号。解决是用sep;重新读入读入后检查shape确认列数是2。第二条Windows下DataLoader多进程报错。现象是进第一个epoch循环就弹RuntimeError提示DataLoader worker异常退出。原因是num_workers大于0时Windows的spawn机制和PIL操作冲突。解决是把num_workers设为0或者把训练主程序包在ifname main:里。这个问题在Windows上几乎必现Linux上反而很少遇到。第三条图像尺寸不一致导致训练中断。现象是某些图片加载后尺寸异常DataLoader里torch.stack报错提示batch内张量尺寸不一致。原因是GTSRB原始图里有少数文件损坏或尺寸特殊Resize没有全部覆盖。解决是在__getitem__里对Image.open加try/except读失败的返回一张同分布的随机噪声图同时打印文件路径。这样训练不会中断你也能定位到具体是哪张图出了问题。第四条过拟合严重训练acc很高但验证只有80%。现象是epoch后半段训练loss持续下降验证loss先降后升。原因有两个数据增强太弱网络容量相对数据量偏大。解决是先加RandomRotation和轻微颜色抖动再把Dropout从0.3提到0.5最后把全连接层512降到256。改完验证acc一般能回升三到五个点。第五条验证指标虚高测试集上掉点。现象是手切验证集后验证acc异常高换测试集低好几个点。原因是没有按原始场景分组同场景连拍泄漏到训练和验证两边。解决是按标志的原始拍摄序列分组切保证同一个指示牌只有一边见过。评估这条线有时候很玄学同样的网络、同样的数据切分方式不同结果能差两个点。所以我现在任何实验都固定切分种子或者直接把csv顺序排好再按比例切保证前后结果可比。血泪经验这条不写进代码里下次实验一定会翻车。6. 把训练好的模型做成单图推理函数一个可直接复用的技巧训练和评估跑通了离能用还差一步——把模型接到真实输入上。我一般会写一个predict_image函数接收任意路径的图片输出类别和概率。这个函数在课程设计答辩、跑小demo、给导师演示的时候直接用不用每次打开训练脚本。from PIL import Image import torch import torch.nn.functional as F device torch.device(cuda if torch.cuda.is_available() else cpu) model TSRCNN(num_classes43).to(device) model.load_state_dict(torch.load(best_model.pt, map_locationdevice)) model.eval() def predict_image(img_path, transformeval_transform): image Image.open(img_path).convert(RGB) tensor transform(image).unsqueeze(0).to(device) with torch.no_grad(): logits model(tensor) probs F.softmax(logits, dim1) pred torch.argmax(probs, dim1).item() return pred, probs[0][pred].item()逻辑说明unsqueeze(0)把单张图变成batch1的4维张量因为网络输入始终是(N,C,H,W)。softmax把logits转成概率分布argmax取最大概率的类别索引。pred是0到42的整数要和csv里的ClassId对应。如果你的标注文件里额外有含义文本比如限速30自己建一个id到中文名的dict输出时直接查表。如果你想把它接到视频流上做实时检测核心还是这个函数——从视频帧里截取标志区域resize到48x48喂给predict_image。注意帧率不要拉太高单张推理在CPU上大约几十毫秒GPU上更快但视频处理真正的瓶颈是检测框怎么找那就是另一个目标检测的话题了。这段代码最常翻车的地方是transform不一致。训练时如果用48x48加归一化推理时也必须用同一套值对不上输出概率全是乱的。我见过有人把这个函数写完后预测结果全是同一个类查了半天是训练和推理用了不同的normalize参数。从那以后我每次换环境都会先打印一张图的预测概率分布看到概率不是集中在某一个类上才敢继续跑。希望帮到你。本文还有配套的精品资源点击获取

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号