ARTICLE DETAIL

资讯详情

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

ResNet图像多分类工程实战:从数据处理到模型训练调优

ResNet图像多分类工程实战:从数据处理到模型训练调优 简介一套面向图像分类初学者的ResNet二维图像多分类完整实现围绕数据准备、残差网络搭建、训练验证与结果可视化展开适合希望快速掌握深度学习分类流程并动手实践的读者。资源共20个文件以Python脚本为主另含编译生成的pyc文件与可视化PNG图压缩包整体约382KB内容覆盖数据预处理、模型定义、训练入口、评估及可视化等环节目录分层清晰。已有215人学习/下载可直接作为入门参考也可修改迁移至自己的图像数据集上训练。通过运行与阅读代码读者能够理解残差连接的基本原理、训练超参数的选择依据并掌握一套通用分类项目的工程组织方式配合数据增强、损失曲线和特征图可视化等实现还能直观体会模型调优与排错思路为后续目标检测、语义分割等更复杂视觉任务打下扎实基础。1. 一个能跑的ResNet 2D图像多分类工程先看清楚它解决了什么拿到一批图像要按内容分类比如产品缺陷分级、聊天截图归类、花卉品种识别手里有时就几千张样本分类数倒是不少。网上的ResNet分类教程大多是单点代码数据处理一段、模型一段、训练一段拼起来各管各的根本不闭环。这份2d_cnn_cls_sample工程就是一个完整的ResNet 2D图像多分类落地包从train_main.py入口进去数据读取、标签生成、增强、训练、进度条、日志、验证全在同一个仓库里接好了跑起来就能看到损失曲线和精度变化。它解决的核心问题是把“一堆散乱图像”变成“一个能用的多分类模型”这条链路上的工程环节补齐而不是卖某个单独算法的新奇度。适合刚入门深度学习、想在真实数据上把ResNet多分类跑通的初学者也适合已经会训练但没空重复搭工具链的工程师。这个方案就是标准的2D卷积分类不做把连续帧2D图像堆叠成3D体数据的那种操作你把自己的数据集按文件夹归类改一下类别数就能跑。2. 目录结构与启动链路train_main.py到utils每一环怎么配合2.1 源文件清单每个文件在链路上的位置整个工程解压之后第一件事不是跑train_main.py而是先确认每个文件在训练链路里处于哪一环。我习惯把工程文件先过一遍避免事到临头发现某个脚本是要单独执行的、某个模块只是被import的库。路径职责在训练链路中的位置train_main.py主入口训练调度执行的起点data_process.py数据集读取、预处理、构建DataLoader数据流上游dataProcess.py训练中的批数据补充处理数据流中游model/resnet.pyResNet模型定义网络构建utils/logger.py日志记录监控utils/bar_utils.py控制台进度条监控utils/misc.py杂项工具函数辅助utils/eval.py验证评估工具训练评估tools.py通用工具集辅助visualize.py训练曲线、样本图可视化分析generateOnehotLabel_txt.py生成one-hot标签txt数据前置imgaug/图像增强相关脚本数据流中游images/示例图像展示data/训练数据目录数据源注意data/和images/不是一回事images/里放的是给别人看的示例图data/下面才是按类别分好的训练图像。你的数据集替换的时候只需要换data/里的内容不要动images/。2.2 启动入口train_main.py训练循环怎么串起来的train_main.py是全局的调度者它做的事可以用一句话概括解析参数、构建数据、建模型、循环训练、保存权重。我挑训练循环的主干部分说明# train_main.py 训练循环核心骨架 import torch import torch.nn as nn from data_process import build_dataloader from model.resnet import resnet18 from utils.logger import get_logger from utils.bar_utils import progress_bar def main(): # 实际工程里这里用argparse接参数比如--data_dir、--num_classes train_loader, val_loader build_dataloader(./data, batch_size32) model resnet18(num_classes10) criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.001, momentum0.9, weight_decay5e-4) for epoch in range(50): model.train() for batch_idx, (images, labels) in enumerate(train_loader): outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() progress_bar(batch_idx, len(train_loader), fEpoch {epoch} loss {loss.item():.4f}) # 每个epoch后调用utils/eval.py里的验证函数逻辑说明train_main.py先通过build_dataloader拿到训练和验证两个DataLoader然后建模型、定义交叉熵损失和SGD优化器。内层循环里每个batch先forward得到outputs再用交叉熵计算loss反向传播后更新参数。进度条由bar_utils控制logger会把每次的loss和accuracy写进日志文件而不是只print到控制台。参数说明batch_size是显存敏感项6G显存一般用32显存大可以提到64显存小就降到16num_classes要和你的数据集类别数量严格一致否则最后一层全连接输出维度对不上SGD的momentum0.9是分类任务常用配置weight_decay5e-4是防止过拟合的正则项2.3 数据读取与工具协同data_process.py、logger.py、bar_utils.py怎么搭的data_process.py负责把data/目录变成模型能吃的Tensor数据。它内部用的是torchvision的datasets.ImageFolder要求data/下面每个子文件夹是一个类别文件夹名就是标签名。比如data/train/dog、data/train/cat读进来之后dog的索引是0cat的索引是1。# data_process.py 数据集构建 import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms def build_dataloader(root, batch_size32, trainTrue): transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) dataset datasets.ImageFolder(root, transformtransform) loader DataLoader(dataset, batch_sizebatch_size, shuffletrain, num_workers2) return loader逻辑说明Resize把所有的图统一到224×224因为ResNet的输入尺寸约定是224ToTensor把HWC的像素数组转成CHW的Tensor同时把像素值从0-255归一化到0-1Normalize再用ImageNet的均值方差做标准化这样和预训练权重的分布对齐。这里的shuffle参数在训练时是True验证时是False顺序打乱是为了避免每个batch里全是同一类图像。logger.py和bar_utils.py是配合项logger把每个epoch的loss、lr、accuracy追加到log文件跑完可以直接看全量趋势bar_utils.py控制台进度条显示当前batch的完成比例和耗时。visualize.py则是在训练结束后画出损失曲线、精度曲线顺便能输出第一批增强后的样本图帮你确认图像没有读坏。3. 模型层实现细节残差块跳跃连接与分类头替换全览3.1 残差块的shortcut为什么从“堆层”改成“跳层”网络加深之后梯度回传时经过多层连乘会越来越小这就是梯度消失。以前的做法是用各种激活函数和归一化去缓解但层数超过一定深度依然难训。ResNet的核心改动是加入了shortcut连接让输入x直接跳到残差块的输出上输出变成F(x)x。这个改动看起来简单但效果很明显梯度可以从最后一层沿着shortcut直接传回前面层即使F(x)里的卷积层梯度很小恒等路径也能保住梯度不回零。所以ResNet可以把网络堆到50层、101层甚至更深而不会像VGG那样到19层就基本训不动。这个工程里用的是2D图像多分类不需要超深网络所以模型文件model/resnet.py里默认实现的是18层结构也就是BasicBlock堆叠的方式。每个BasicBlock包含两个3×3卷积中间用BN和ReLU隔开外加一个shortcut。3.2 resnet.py里的BasicBlock与ResNet主体model/resnet.py是理解整个模型的关键文件。看BasicBlock的代码时重点看两个地方shortcut什么时候需要加1×1卷积两个3×3卷积的stride分别是什么。# model/resnet.py BasicBlock核心实现 import torch.nn as nn class BasicBlock(nn.Module): expansion 1 def __init__(self, in_channels, out_channels, stride1): super().__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm2d(out_channels) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm2d(out_channels) ) def forward(self, x): identity x out self.conv1(x) out self.bn1(out) out self.relu(out) out self.conv2(out) out self.bn2(out) out self.shortcut(identity) out self.relu(out) return out逻辑说明forward里先保存identity为输入x然后依次过conv1、bn1、relu、conv2、bn2最后把shortcut的结果加在out上再经过ReLU输出。当stride1且输入通道等于输出通道时self.shortcut是空的等于恒等连接当stride2或者通道数变化时shortcut用1×1卷积加BN把输入调整到和输出同样的尺寸和通道数。参数说明biasFalse是因为后面接BatchNorm卷积层不需要额外的偏置参数padding1保证3×3卷积在stride1时特征图尺寸不变inplaceTrue让ReLU直接修改原Tensor省显存3.3 下采样策略stride2怎么让特征图减半ResNet的4个stage分别对应不同分辨率的特征图。输入224×224的图像经过7×7卷积和3×3池化后变成56×56再经过layer1到layer4最终特征图变成14×14。整个过程靠的就是每个stage第一个BasicBlock的conv1的stride2。Stage输入尺寸输出尺寸通道数关键操作conv1224112647×7 stride2layer111211264BasicBlock stride1layer211256128第一个块stride2layer35628256第一个块stride2layer42814512第一个块stride2这个规律在每个stage里都是固定的只在该stage的第一个残差块把stride设成2其余残差块保持stride1。特征图减半的同时通道数从64翻到128、256、512相当于用空间分辨率的降低换取更丰富的语义特征。这也是为什么类别靠后的层反而“看得更抽象”。3.4 分类头替换avg_pool fc变成num_classes之前的层都在做特征提取最后的全局平均池化和全连接层才是真正的分类头。原始ResNet在ImageNet上输出1000类你的任务如果是10类就要把最后一层Linear的out_features改成10。# model/resnet.py ResNet主体 import torch.nn as nn class ResNet(nn.Module): def __init__(self, block, num_classes10): super().__init__() self.conv1 nn.Conv2d(3, 64, kernel_size7, stride2, padding3, biasFalse) self.bn1 nn.BatchNorm2d(64) self.relu nn.ReLU(inplaceTrue) self.maxpool nn.MaxPool2d(kernel_size3, stride2, padding1) self.layer1 self._make_layer(block, 64, 64, 2, stride1) self.layer2 self._make_layer(block, 64, 128, 2, stride2) self.layer3 self._make_layer(block, 128, 256, 2, stride2) self.layer4 self._make_layer(block, 256, 512, 2, stride2) self.avgpool nn.AdaptiveAvgPool2d((1, 1)) self.fc nn.Linear(512 * block.expansion, num_classes) def _make_layer(self, block, in_channels, out_channels, blocks, stride): layers [] layers.append(block(in_channels, out_channels, stride)) for _ in range(1, blocks): layers.append(block(out_channels, out_channels, stride1)) return nn.Sequential(*layers)逻辑说明_make_layer的第一个块传入stride参数这就是之前说过的下采样入口后续块保持stride1。avgpool用AdaptiveAvgPool2d把任意尺寸的特征图池化成1×1然后展平后接全连接层。最后一个Linear的输出维度就是num_classes。这里有个值得注意的点AdaptiveAvgPool2d和标准AvgPool2d不同它不需要你手动算池化核大小输出尺寸固定为1×1省去了一堆维度对齐的麻烦。3.5 resnet预训练模型的加载权重怎么接、哪层会被跳过自己随机初始化训练不是不行但在数据量只有几千张时容易欠拟合。常见做法是加载在ImageNet上预训练好的resnet18权重让模型先在通用特征上起步再在你的分类任务上微调。工程里model/resnet.py可以自由切换为预训练版本# 加载resnet预训练模型的常见做法 import torchvision.models as models import torch.nn as nn def get_resnet(num_classes, pretrainedTrue): model models.resnet18(pretrainedpretrained) in_features model.fc.in_features model.fc nn.Linear(in_features, num_classes) nn.init.kaiming_normal_(model.fc.weight) return model逻辑说明先把官方预训练权重读进来然后拿到fc层原本的输入维度in_features替换成新的全连接层。因为你自己任务里的类别数不是1000所以最后一层的权重参数形状对不上直接load的话会报错。这里先替换fc再用kaiming_normal_对新fc层做初始化。提示如果加载模型后直接调用model.load_state_dict带预训练权重的checkpoint会报Missing key(s) in state_dict这是正常的因为fc层的key不匹配。在工程里我一般用strictFalse加载然后把fc层单独初始化这样既保留了前面4个stage的特征提取能力又让分类头从零开始学。4. 数据准备与增强one-hot标签生成、预处理与划分避坑4.1 generateOnehotLabel_txt.pyone-hot标签文件怎么生成在训练之前数据准备是绕不开的一步也是最容易出细节问题的一步。工程里提供了generateOnehotLabel_txt.py它做的事情是扫描data目录下的所有图片把图片路径和one-hot标签写进一个txt训练时再按这个txt逐行读取。为什么需要这个脚本因为ImageFolder虽然能按子目录名自动生成标签但它不保证多次运行之间文件顺序完全一致。当你需要固定训练集和验证集划分时先导出一份带路径和标签的txt相当于给数据快照上了一把锁。# generateOnehotLabel_txt.py 生成one-hot标签文本 import os def generate(label_file, class_names, image_dir): with open(label_file, w, encodingutf-8) as f: for root, _, files in os.walk(image_dir): for name in files: if not name.lower().endswith((.jpg, .jpeg, .png)): continue path os.path.join(root, name) label os.path.basename(root) idx class_names.index(label) onehot [0] * len(class_names) onehot[idx] 1 f.write(f{path}\t{ .join(map(str, onehot))}\n)逻辑说明脚本用os.walk遍历data目录每一张图片的路径是完整路径类别名取自它的上级目录名class_names是预先定义好的类别列表。one-hot向量长度等于类别总数只在当前类别索引位置写1。生成的每一行是“图片路径\t0 1 0 0”这样的格式。参数说明class_names的顺序必须和后续训练时的类别列表保持一致否则标签就是错的只筛选.jpg、.jpeg、.png三种常见后缀避免隐藏文件混进来用完整路径比相对路径更稳妥因为后续训练脚本可能在别的目录下启动生成之后我习惯用文本编辑器打开txt抽查第一行、中间一行、最后一行确认路径存在、标签对应正确。这一步十分钟能做完但能省下后面排查标签错位的几天时间。4.2 图像预处理resize到224、归一化和通道顺序图像读进来是H×W×C的格式而PyTorch的卷积层要求N×C×H×WToTensor()会顺带完成这个转置。但除了通道顺序还有两个参数必须注意。第一个是输入尺寸。ResNet的预训练权重是按224×224尺寸训出来的如果你的图像是500×500直接丢给模型一旦后面接全连接层维度会算不对。工程里统一Resize到224×224这是各类分类网络最通用的尺寸。第二个是归一化。单张图像的像素值范围是0-255ToTensor()会除255变成0-1之间但网络里的BN层期望的是接近标准正态分布的输入。所以train_main.py那条链路里用了transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])这组mean和std是ImageNet百万张图的统计值加载预训练权重时用它们能让输入分布和预训练时保持一致。如果你不用预训练权重可以统计自己数据集的均值和方差但绝大多数场景直接沿用这组数没有大问题。4.3 数据集划分类别均衡的分层切分训练集、验证集、测试集的划分要处理好不然会出现训练精度差、验证精度虚高、或者某个类别在验证集里完全没有样本的情况。最常见也最容易被忽略的是分层划分问题。假设你有10个类别每个类别的样本数量不同如果不做stratify直接随机切分那些样本少的类别可能全部落在训练集里验证集里看不见它。训练时模型对这种类别的识别能力完全是黑匣子。用sklearn的train_test_split做分层划分是常见做法from sklearn.model_selection import train_test_split def split_by_files(file_list, labels, val_ratio0.1, seed42): train_files, val_files train_test_split( file_list, test_sizeval_ratio, stratifylabels, random_stateseed) return train_files, val_files逻辑说明stratifylabels让切分后训练集和验证集里各个类别的比例与原数据集一致random_state固定随机种子保证两次划分结果相同。这样无论是复现还是后续调参数据划分的基线是一致的。划分比例上我一般用90%训练、10%验证如果数据量少于2000张可以改成85%训练、15%验证。验证集不是为了训练而是为了在每个epoch后检查模型是否过拟合所以比例不要太少。4.4 imgaug增强小数据集撑出多样本的正确姿势数据增强是防止过拟合最直接的手段。工程里imgaug/目录封装了图像增强的配置它基于imgaug库实现和你平时用的PyTorch RandomHorizontalFlip不同它可以做更丰富的几何和颜色扰动。# imgaug/ 增强管线示意 import imgaug.augmenters as iaa seq iaa.Sequential([ iaa.Fliplr(0.5), iaa.Affine(rotate(-10, 10), scale(0.9, 1.1)), iaa.Multiply((0.9, 1.1)), iaa.AdditiveGaussianNoise(scale(0, 0.05 * 255)) ])逻辑说明这条管线定义了50%概率水平翻转、±10度随机旋转、0.9到1.1倍的随机缩放、亮度乘0.9到1.1、以及少量高斯噪声。管线会在训练时对每一张图随机应用这些变换相当于每一轮epoch喂给模型的数据都不完全一样。参数说明rotate的范围不要超过±15度角度太大会破坏图像的语义比如把数字6旋转后看起来像9scale的0.9到1.1表示缩放10%以内过大会导致关键部位被裁掉AdditiveGaussianNoise的scale值不要超过0.05×255否则图像会明显发花增强只加在训练集上验证集和测试集只用Resize和Normalize不能把随机增强加在验证集上否则验证精度会忽高忽低失去参考价值。5. 训练调和与避坑排查五个高频问题的现象、原因、解决5.1 训练主参数优化器、学习率与学习率调度怎么配训练阶段的参数配置直接决定模型能不能收敛。工程里默认用的是SGD加momentum损失函数用交叉熵。这个组合在图像分类任务上比Adam更稳Adam前期收敛快但后期精度容易被SGDmomentum反超。criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr0.001, momentum0.9, weight_decay5e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size20, gamma0.1)逻辑说明CrossEntropyLoss内部自带softmax不需要在网络输出层再加。StepLR每20个epoch把学习率乘以0.1这样训练前中期用较大学习率快速靠近最优区域后期用较小学习率精细收敛。参数说明lr0.001在加载预训练权重时是安全的起点如果模型随机初始化可以试着调到0.01weight_decay5e-4是L2正则防止大权重把训练集特征背下来scheduler的step_size要和总epoch数匹配50个epoch时20一轮比较合适100个epoch时可以改成30这几项配好之后不需要频繁动真正决定模型上限的往往是数据和标签质量。5.2 踩坑一训练loss卡在平台期不降现象loss前几个epoch降得很顺利降到0.3左右就再也不动了验证精度也停在60%附近上不去。原因最常见的是学习率太小。SGD在0.001的学习率下到了loss比较平坦的区域后更新量不够参数走不动。第二个原因是batch_size太小比如只有4或8BN的batch统计不稳定导致loss曲线上下抖但均值不动。解决先把学习率调高一个数量级试试如果是0.001就改成0.005或0.01观察5个epoch看loss是否降。如果调高后loss发散说明原来的0.001已经是上限问题就不在学习率而是在数据增强太强或者标签有错。另外把batch_size提到16以上BN的统计才会稳定。5.3 踩坑二训练精度很高、验证精度上不去现象训练集精度到96%验证集只有72%每轮epoch验证集精度都不怎么涨。原因典型的过拟合。模型容量对当前数据量来说太大了或者增强太弱让模型直接把训练样本背了下来。还有一个隐蔽原因数据划分时没有做分层某个类别在训练集里出现特别多在验证集里出现特别少模型对那个类别的泛化能力自然差。解决训练集增强先开满包括翻转、旋转、亮度扰动确认weight_decay是开着的5e-4起步如果用了预训练权重可以冻结前面stage只训练最后的分类头等方法起效再解冻微调。最后检查划分逻辑用stratify保证每个类别在验证集里的比例和训练集一致。5.4 踩坑三训练中途CUDA out of memory现象前几个epoch跑得好好的跑到第8个epoch突然报显存OOM程序退出。原因多数不是显存真的被模型占满而是验证阶段没有包在torch.no_grad()里验证时还在构建autograd计算图。另一个常见原因是验证集的transform里忘了Resize原始大图直接喂进网络一个batch的内存占用瞬间暴涨。解决把所有验证和测试代码都放进with torch.no_grad()块里确认验证集的transform和训练集一样都包含Resize((224,224))。如果batch_size32跑不过降到16或者8配合梯度累积也能保持相同的更新步长。还有一个技巧是在验证循环结束后调用torch.cuda.empty_cache()释放显存碎片。5.5 踩坑四加载权重报Missing key与Unexpected key现象训练完保存了模型下次想加载继续训练model.load_state_dict直接报key名对不上。原因两类最常见的情况。第一种是训练时用了DataParallel它会把模型参数包一层module保存下来的key都是module.conv1.weight这种前缀而当前模型是裸的key对应不上。第二种是你加载了ImageNet预训练权重但自己的fc层换过维度原始的1000类全连接key自然不存在。解决如果是DataParallel的前缀问题加载前把key里的module.去掉state_dict torch.load(best.pth) new_state_dict {k.replace(module., ): v for k, v in state_dict.items()} model.load_state_dict(new_state_dict)如果是fc层不匹配用model.load_state_dict(state_dict, strictFalse)加载之后重新初始化fc层。5.6 踩坑五标签错位导致loss跳来跳去现象loss曲线上下剧烈跳动从0.8跳到1.2再跳回0.7训练精度一直在30%-40%随机水平怎么调学习率都没用。原因这类情况十有八九是标签和图像对不齐。自己写了generateOnehotLabel_txt.py之后txt里的路径顺序和DataLoader读入的顺序不一致或者是做数据划分时只用了一个文件列表但标签列表没有跟着一起切分导致图像是狗的图标签却写成了猫。解决训练脚本里完全依赖txt中记录的路径读取图像不要用两个独立的列表分别放路径和标签。每次换新数据集时把第一个batch的数据打印出来用matplotlib画成九宫格每张图标题显示它的标签索引人工看一遍。这一步看似原始但标签错位是最难排查的错误靠看loss曲线是看不出原因的。6. 迁移到自己的数据集改五分类、混淆矩阵与阈值过滤6.1 把三分类改成五分类需要动哪几处工程默认的num_classes是10改成你的任务只需要动三处data/目录下的子文件夹数量、train_main.py的--num_classes参数、以及generateOnehotLabel_txt.py里class_names列表。如果你的任务是五分类class_names写五个文件夹名num_classes传5其他都不用改。注意如果之前跑过训练旧的日志文件和模型权重要先删掉或者用新路径命名否则加载模型时可能把之前5类的模型当成预训练权重用。6.2 用python多分类混淆矩阵代码做细粒度评估训练结束后整体的accuracy只能告诉你大概水准无法告诉你哪些类别在互相混淆。这时要看混淆矩阵。from sklearn.metrics import confusion_matrix, classification_report import matplotlib.pyplot as plt import seaborn as sns y_true, y_pred [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: outputs model(images) _, preds torch.max(outputs, 1) y_true.extend(labels.numpy()) y_pred.extend(preds.numpy()) cm confusion_matrix(y_true, y_pred) sns.heatmap(cm, annotTrue, fmtd, cmapBlues) plt.xlabel(Predicted) plt.ylabel(True) plt.show()逻辑说明模型在验证集上逐batch预测preds是每个样本得分最高的类别索引。最后用sklearn的confusion_matrix生成矩阵seaborn画出热力图对角线越亮说明这个类别的准确率越高。非对角线上的亮点就是容易互相误判的类别对。只看准确率会漏掉一个关键信息如果类别A总被预测成类别B说明这两类图像在视觉上或者拍摄条件下太接近了。此时可以回看训练集里这两类的样本判断是标注本身模糊还是模型特征区分力不够。6.3 预测完成后加一道置信度后处理多分类任务里torch.max取概率最高的类别通常就够了但工程场景往往需要拒绝不可靠的预测。比如置信度0.35就把它当成最终结果在质检场景里会埋雷。我的习惯是在输出层加一道置信度判断低于阈值就归为“不确定”probs torch.softmax(outputs, dim1) max_prob, pred torch.max(probs, dim1) final torch.where(max_prob 0.8, pred, torch.tensor(-1))阈值0.8可以在验证集上调观察的是阈值提高会牺牲多少召回率换来多少误检下降。这个取舍在工业应用里比单纯堆模型精度更重要。从那以后我每次换数据集都强制走一遍这三步生成标签后打印三处随机行的路径与标签训练5个epoch后人工检查损失曲线和第一批batch的图像与标签对最后用混淆矩阵验证而不是只看准确率。这三步最多花半小时但能拦住标签错位、过拟合和类别混淆三类大坑。希望帮到你。本文还有配套的精品资源点击获取
返回列表