尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

EfficientNet迁移学习实战:从预训练权重到自定义数据集训练全解析

EfficientNet迁移学习实战:从预训练权重到自定义数据集训练全解析 简介这是基于EfficientNet-PyTorch框架训练自定义图像分类数据集的完整演示资源面向正在做分类任务、想借助迁移学习快速出效果的开发者。资源以简洁工程的方式展示了从数据准备到训练测试的完整路径数据集按train和test目录组织每个类别一个子文件夹代码中预留了调整类别数、图像尺寸、批大小等核心参数的位置并支持选择是否自动下载预训练权重便于在自定义数据上微调从而提升小数据集上的收敛速度与准确率。压缩包仅10KB共5个文件其中4个Python脚本构成训练、测试与工具模块1个Markdown文档说明使用要点体量小巧却覆盖关键流程。目前已有3780人学习参考可视为图像分类入门迁移学习的实用范本拿到后只需替换自己的数据目录并修改对应参数即可运行能省去大量环境与代码调试时间。1. 在 EfficientNet 上训练自己的数据集这份源码包怎么用用 EfficientNet 训练自己的分类数据集最容易被卡住的其实不是模型结构而是数据目录和训练脚本之间的耦合。EfficientNet-Pytorch 这个演示项目正好把这条路径打通了预训练权重下载、数据集读取、训练测试循环全部收敛在一个 efficientnet_sample.py 里改几行就能从 ImageNet 权重迁移到自己的任务上。适合手里有自定义分类数据但不想从零写训练框架的从业者也适合第一次尝试 EfficientNet 系列 backbone 的人。先跑通再慢慢改参数这套流程能让你用最短时间看到模型在自己的数据上的真实水平。2. EfficientNet 选型与数据集组织目录结构定成败2.1 复合缩放与迁移学习为什么偏向 EfficientNet-B0 起步图像分类任务里选 backbone大家通常会纠结 ResNet、VGG 还是 EfficientNet。VGG 结构简单好理解调试起来很直观但参数冗余太大训练和推理都吃亏ResNet 生态成熟改起来资料多但同等精度下参数量和计算量都不占优势。EfficientNet 走的是另一条路它先用神经网络搜索出一个高效的基线网络 MBConvBlock然后通过复合缩放方法同时把网络的深度、宽度和输入分辨率按比例放大。也就是说它不是单纯地把每一层变宽或者堆更多层而是三个维度一起控制这样在相同计算量预算下能拿到更高的精度。以默认的 efficientnet-b0 为例它的参数量大约是 5.3M输入分辨率是 224×224在 ImageNet 上能做到约 77% 左右的 top-1 精度。这个配置放在迁移学习场景里非常合适因为它足够轻单卡可以跑训练速度快而且预训练权重已经覆盖了通用视觉特征。我在实际项目中一般建议从 B0 起步先把整个流程跑通确认数据没问题之后再考虑换 B2 甚至 B3 提精度。如果一上来就选 B5显存压力大不说训练速度慢还会让参数调优变得很痛苦。选型还有一个容易被忽略的点EfficientNet 的预训练权重和模型结构是强绑定的不同变体的 depth 系数和 width 系数都不同你不能用 B0 的权重去初始化 B2 的网络尺寸根本不匹配。所以换模型变体就要重新准备对应的预训练权重这个后面会讲。整体来说对于自定义分类数据集B0 是非常妥当的起点尤其当你的数据量在一万到十万张这个区间时B0 配合迁移学习已经能拿到很不错的结果。2.2 数据集目录怎么摆train/test 下的类别子文件夹是关键这个项目的数据读取方式依赖 torchvision 的 ImageFolder 机制它要求数据集按类别放在子目录里。源码包没有自带图片数据需要你自己把数据集整理成这样的形式dataset/ train/ 1/ image_0001.jpg image_0002.jpg ... 2/ ... test/ 1/ 2/ ...这里的 1、2 就是类别编号脚本训练时会把文件夹名映射成类别索引。比如 train/1 下的图片全部对应标签 0train/2 下的图片对应标签 1。映射顺序取决于文件夹名的排序方式这一点要特别留意。我在组织自己的数据集时习惯把数字编号换成有意义的类别名比如 dog、cat这样在分析错误样本时不用再去翻标签表。ImageFolder 读取时还有一个特性它会把 test 目录也按同样的方式映射一遍。所以 train 和 test 的子文件夹名必须保持一致否则类别索引就会错位。比如 train 下面有 1、2、3 三个文件夹test 下面却只有 1、2训练时 num_classes 设为 3测试时实际只有两类准确率会很低且不具参考性。关于类别数量要特别提醒一下。EfficientNet 默认的 num_classes 是 1000脚本里提供给我们改的是第 15 行左右的 num_classes 参数。如果你只改了类别数但没改训练数据目录结构脚本运行时会因为找不到类别文件夹直接抛异常。我在自己的一个缺陷分类项目里一开始就漏了这个结果跑了半天才发现是数据目录没建好白费了很多时间。2.3 增强策略与小样本处理翻转裁剪之外还有重采样自定义数据集最常见的痛点是数据量不够尤其某些类别可能只有几十张图。EfficientNet 在迁移学习中遇到小数据集过拟合会很严重。这个演示脚本本身没有做复杂的增强只是做了基本的随机裁剪和翻转但我们在实际使用的时候完全可以自己在数据加载环节补上更合适的增强配置。我一般会这样写# 在数据加载部分加入增强 data_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.8, 1.0)), transforms.RandomHorizontalFlip(p0.5), # 仅当翻转不改变语义时使用 transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])这段增强配置中RandomResizedCrop 的作用是随机裁剪并缩放到 224×224scale 参数控制在 0.8 到 1.0 之间避免裁掉太多主体。RandomHorizontalFlip 适合那些左右翻转不会改变语义的任务比如猫狗分类但如果你做的是数字识别或者左右不对称的物体分类这个操作会引入错误标注。ColorJitter 调整亮度和对比度对光照变化较大的真实场景数据很有帮助。归一化的均值和标准差直接用 ImageNet 的默认值因为迁移学习的 backbone 已经习惯了这套分布。除了增强类别不平衡也要在数据准备阶段处理。我在一个工业缺陷检测项目里正负样本比例接近 30 比 1直接训练出来的模型对少数类几乎没有召回。后来做了类别重采样保证每个 epoch 里各个类别出现的样本数接近效果才正常。实现方式也比较简单可以用 WeightedRandomSampler 或者直接对少数类做上采样都能缓解这个问题。3. 环境配置与 efficientnet_sample.py 关键代码拆解3.1 先搭一个能跑的老环境PyTorch 版本怎么选拿到这个项目的第一步是给 efficientnet_pytorch 模块找到能跑的 PyTorch 环境。它的依赖很轻核心就是 torch 和 torchvision外加 Pillow 读图。我自己通常会用 Anaconda 创建一个独立环境避免污染其他项目conda create -n effnet python3.8 conda activate effnet conda install pytorch torchvision cudatoolkit11.3 -c pytorchPython 版本建议控制在 3.8 左右。EfficientNet-Pytorch 是较早的代码库在 Python 3.10 以上环境中导入时偶尔会遇到 typing 相关的兼容问题。PyTorch 版本从 1.10 到 2.x 都能用但需要强调的是torchvision 的版本必须和 torch 严格对应否则 import 阶段就会报错。常见对应关系可以参考下面这张表PyTorch 版本torchvision 版本Python 兼容性1.100.113.6-3.91.120.133.7-3.102.00.153.8-3.112.10.163.8-3.11如果没有 GPUCPU 版本也能把流程跑通只是速度会比较慢。B0 在 CPU 上训练 224 分辨率的图一个 batch 大约需要零点几秒一个 epoch 几百张图也要跑上几分钟。验证代码逻辑没问题但要训练到收敛还是建议用 GPU。这个资源包里的文件结构很清晰核心模块在 efficientnet_pytorch 目录下文件作用model.pyEfficientNet 网络结构定义utils.py预训练权重加载、模型构造辅助函数init.py包导出入口efficientnet_sample.py训练与测试入口脚本README.md原作者的说明3.2 line 13-22 参数区逐行解读打开 efficientnet_sample.py最先要动的是顶部这块参数区。默认的配置是针对演示数据集写的必须改成你自己的数据集配置。这段代码我把它搬出来逐行说明# line 13-22 参数区 num_epochs 60 # 训练轮数数据量大时适当减少 batch_size 16 # 批量大小显存不足时降到 8 num_classes 5 # 类别数必须和数据集子文件夹数一致 lr 0.01 # 初始学习率迁移学习建议降到 0.005 以内 momentum 0.9 # SGD 动量保持默认即可 weight_decay 5e-4 # 权重衰减防止过拟合的重要参数 step_size 30 # 学习率每多少轮衰减一次 gamma 0.1 # 衰减系数 model_name efficientnet-b0 # 可选 efficientnet-b0 到 b7第一个必须改的是 num_classes它要和你的数据集类别数一致。第二个是 model_name默认 b0如果你的显存足够而且数据量很大改成 b2 能明显提精度。batch_size 的设定要看显存16G 显存跑 b0 用 16 没问题8G 显存建议降到 8否则容易 OOM。学习率 0.01 对从头训练一个网络来说是合理值但迁移学习场景下这个值偏大。预训练权重已经包含了很多通用特征过大的学习率会破坏这些特征导致前期 loss 震荡甚至发散。我在自己的项目里一般从 0.001 开始调如果 loss 下降太慢再逐级上调。momentum 和 weight_decay 是 SGD 的常规配置不熟悉就先不要动。3.3 预训练权重下载逻辑与 eff_weights 目录line 169 附近是预训练权重的控制逻辑主要作用是决定加载 ImageNet 预训练权重还是从随机初始化开始。默认设置为自动下载脚本会先检查本地 eff_weights 目录有就直接加载没有就尝试通过 url 下载# line 169 附近 load_pretrained_weights True # True 下载并加载预训练权重False 从头训练这里有一个很影响进度的问题代码默认从外网下载权重如果你的训练环境无法访问外网脚本不会立刻报错而是长时间卡在下载阶段。我在内网服务器上踩过这个坑当时还以为是训练卡住了等了十几分钟才发现是在下载。解决方法是先在一台能联网的机器上把权重下好然后传到 eff_weights 目录下。权重文件的命名必须和代码里的拼接逻辑一致比如 efficientnet-b0 对应的权重文件名通常是 efficientnet-b0.pth不要自己随意改名。如果需要离线部署我建议把整个 eff_weights 目录连同代码一起打包分发这样目标机器上就不需要任何外网访问能力。3.4 model.py 与 utils.pybackbone 封装的黑匣子怎么用训练脚本通过 efficientnet_pytorch 包导入模型核心实现都在 model.py 里。这个文件里定义了 EfficientNet 的完整结构包括 MBConvBlock、SE 模块和最终的分类头。我们在使用时通常不需要修改它只需要通过 from_pretrained 方法加载权重# 构造模型并加载预训练权重 from efficientnet_pytorch import EfficientNet model EfficientNet.from_pretrained(efficientnet-b0, num_classesnum_classes)from_pretrained 是 utils.py 里提供的便捷方法它会根据模型名自动找到对应的预训练权重文件并加载同时把最后一层分类头替换成你传入的 num_classes 数量。这里要注意一点如果模型名写错了比如把 efficientnet-b0 写成 efficientnet_b0它会尝试下载一个不存在的权重文件报错信息是 HTTP 404看起来像网络问题。此外model.py 还有一个针对灰度图的处理逻辑。如果你的输入是单通道灰度图它会把第一个卷积层的输入通道从 3 改成 1但这个改动和预训练权重是不匹配的。我在实际项目中吃过这个亏灰度图转成 RGB 三通道之后直接把三个通道都填入相同的灰度值这样既保留了预训练权重的完整性效果也更好。4. 训练与测试实操从第一个 epoch 到 acc 稳定4.1 数据丢进去之前先把这四种问题检查掉执行训练前我会先做一遍数据体检把最容易翻车的几个点全部排掉。第一个是检查图片是否损坏比如下载的数据集中有些文件其实不是图片PIL 打开时会抛异常ImageFolder 在读取时会直接中断训练。处理方式是在数据加载前写一个快速脚本把所有图片用 PIL 打开一遍打不开的直接删除或者移到临时目录。第二个是检查类别文件夹名是否和标签对应。train 和 test 下的子文件夹名必须一致否则类别索引错位测试结果会非常离谱。第三个是检查 num_classes 是否真的和文件夹数一致。我曾经因为数据集里有空文件夹导致 ImageFolder 实际读到的类别数比我自己设的少了一个训练时 size mismatch 报错排查了半天。第四个是检查 train 和 test 的数据分布尽量保证每个类别在测试集中都有覆盖否则会有一个类别完全无法验证。# 快速统计 train 下每个类别的图片数量 find dataset/train -type f -name *.jpg | awk -F/ {print $3} | sort | uniq -c这条命令会把 train 目录下每个别的图片数量统计出来一眼就能看出有没有类别是空文件夹或者数据量严重不平衡。4.2 训练循环的观察点loss 不降先别动网络启动训练的方法很简单环境配好、参数改好、数据放好之后直接执行python efficientnet_sample.py脚本会自动完成数据加载、模型构建和训练循环。每个 epoch 结束后终端会打印 loss 和 acc我要盯着看的核心指标有三个训练 loss 是否持续下降、测试 acc 是否稳步上升、训练 acc 和测试 acc 的差值是否急剧扩大。如果训练 loss 下降但测试 acc 停在某个值不动这是典型的过拟合信号优先去调 weight_decay 和数据增强而不是换更复杂的网络。学习率衰减策略这里值得多说一句。默认的 step_size30 和 gamma0.1 意味着在第 30 轮时学习率乘 0.1总共 60 轮的训练只衰减一次。我实际用下来把 step_size 改成 20 会更顺畅相当于在第 20 轮和 40 轮各衰减一次最后几轮的 loss 曲线会更平滑。如果数据集比较小、训练轮数不需要 60 轮那么 step_size 也要跟着缩原则是让它刚好落在训练总轮数的三分之一到二分之一之间。我在实际训练中还会改一个东西每轮训练后保存一份 checkpoint 到 dataset/model 目录。脚本默认只保存最终模型如果训练到一半进程崩溃前面所有的训练时间就全浪费了。加几行逻辑每个 epoch 保存一次看起来简单但在长训练里作用很大。4.3 测试阶段怎么验证只看 acc 会漏掉一种陷阱训练完成后脚本会在 test 集上算整体准确率但这个数字只是一个综合指标它掩盖了不同类别之间的巨大差异。我每次都坚持做逐类别的统计方法是在训练完加载最佳权重后遍历 test 目录做推理# 加载训练好的最佳权重逐类别计算准确率 import torch from efficientnet_pytorch import EfficientNet model EfficientNet.from_pretrained(efficientnet-b0, num_classesnum_classes) model.load_state_dict(torch.load(dataset/model/best_model.pth)) model.eval() class_correct {} class_total {} for class_name in test_class_names: class_correct[class_name] 0 class_total[class_name] 0 with torch.no_grad(): for image_path, true_label in test_samples: outputs model(transform(image_path).unsqueeze(0)) _, predicted torch.max(outputs, 1) class_correct[class_name] (predicted.item() true_label) class_total[class_name] 1这段代码的核心是遍历所有测试样本把预测结果和真实标签一一对比然后按类别累加。加这个统计之后你会发现很多隐藏在整体 acc 背后的问题比如某个类别只有 60% 准确率其他类别都在 95% 以上。这种分类通常代表训练数据里该类别与其他类别区分度不够或者训练样本太少后续针对性地补充数据比调整网络结构更有效。也可以顺手把错误样本的路径打印出来直观地看是标注错了、图片模糊还是模型确实分不清这个类别。这一步在真实项目里能省下很多盲目调参的时间。4.4 模型保存与恢复model 目录下的 best_model.pth脚本会把模型权重保存在 dataset/model 目录。如果你的工作目录结构和标准结构不一致记得先手动创建 dataset/model 文件夹否则 torch.save 会报目录不存在的错误。加载恢复时也要注意两点第一是构造模型时传入的 num_classes 必须和训练时一致否则 load_state_dict 会报 size mismatch第二是加载前调用 model.eval() 切换到推理模式这会关闭 dropout 和 batch norm 的统计更新保证预测结果稳定。# 推理时加载模型的正确姿势 model EfficientNet.from_pretrained(efficientnet-b0, num_classesnum_classes) state_dict torch.load(dataset/model/best_model.pth, map_locationcpu) model.load_state_dict(state_dict) model.eval()map_locationcpu 是为了在没有 GPU 的机器上也能加载权重。如果你的训练环境和推理环境 GPU 型号不同加载时可能会碰到键值对不匹配的问题这个参数可以最大程度避免它。5. 踩坑记录训练自数据集时最容易翻车的五个点5.1 现象运行脚本后卡在下载预训练权重进度条长时间不动如果把训练脚本放在内网服务器上启动最常见的现象是终端卡住没有报错也没有进展。原因是脚本默认尝试从外网拉取预训练权重而内网环境无法访问这个地址socket 一直处于等待状态。解决方法是先在一台能联网的机器上下载好对应模型文件然后通过 scp 传到 eff_weights 目录文件名必须和代码里拼接出来的完全一致。脚本启动时会优先检查本地目录文件存在就跳过下载。5.2 现象训练第一轮就报 size mismatch提示 fc.weight 形状对不上这个错误几乎每个用过迁移学习的人都会遇到。原因很简单预训练模型的分类头是 1000 类你的数据集可能只有 5 类直接加载权重时最后一层的 weight 和 bias 形状不匹配。这不是代码 bug而是正常的迁移学习行为。解决方法是保证 from_pretrained 调用时传入 num_classes 参数这个参数会把分类头替换成你设定的数量并在加载预训练权重时自动跳过分类头这一层。5.3 现象训练过程中显存不足OOM 报错B0 虽然轻量但如果你在复现时改大了 batch_size 或者输入的图片分辨率没有缩放到 224显存会快速耗尽。我见过有人直接把原始 3000×2000 的图喂进去显存瞬间爆掉。解决方法是先检查输入预处理有没有 Resize 或 RandomResizedCrop把分辨率限制在 224 或 260 附近然后把 batch_size 降到 8 或者 4 重试。如果还想用更大的模型或者更高的分辨率可以考虑梯度累积把多个 batch 的梯度累计后再更新参数。5.4 现象训练 loss 一直降不下去徘徊在某个值附近loss 不降通常有两个原因。第一个是学习率设得过高梯度在最优解附近震荡loss 在小范围不断跳动。这种情况的特征是 loss 曲线很毛糙数值没有明显下降趋势。解决方法是把 lr 降一个数量级比如从 0.01 降到 0.001。第二个原因是数据本身噪声大比如标签标错了或者同一类别的图片在语义上差异太大。这种情况需要去人工检查训练样本找到并修正错误标注比调任何参数都有效。5.5 现象跑同样的代码两次训练结果差很多如果数据量较小随机初始化会对结果产生不小影响。同一个脚本跑两次测试 acc 的波动可能达到好几个百分点。这种情况不是代码问题而是数据集本身对随机性敏感。解决方法是固定随机种子包括 torch 和 numpy 的种子让每次训练结果可复现。更进一步的做法是训练三次每次用不同的随机种子最后对三个模型的预测做投票稳定性会明显提升。5.6 现象类别标签错位train 和 test 的文件夹名不一致ImageFolder 会按字典序给类别编号如果 train 下是apple, banana, cherrytest 下是cherry, apple, banana映射关系完全不一致测试准确率会低得离谱。这个问题的隐蔽之处在于代码不报错只是结果很差。解决方法是写一个脚本对比 train 和 test 的文件夹名列表确保完全一致这个检查我应该早点做那次的教训挺深刻。5.7 现象torch.load 在加载模型时提示键值对不匹配保存模型时可能用了多 GPU 并行权重键名里多了module.前缀加载到单卡模型时会匹配不上。或者保存时的 num_classes 和加载时不一致也会报这个错。解决方法是加载时打印 state_dict 的 keys对比一下模型构造器的 keys做成映射关系去除前缀保证两侧结构完全对齐。6. 进阶改造把演示脚本升级成可监控的分类训练器6.1 用 TensorBoard 替代终端日志单纯看终端打印的 loss 和 acc很难感知训练的全局趋势。我会在脚本里接入 TensorBoard每个 epoch 写入训练 loss、训练 acc、测试 acc 和当前学习率训练结束后在浏览器里直接看得一清二楚。改造成本很低只需要在训练循环里加两行 SummaryWriter 调用换来的是对模型训练状态的完整掌控。6.2 早停机制与最佳权重回滚当验证集 acc 连续多个 epoch 没有提升时继续训练只会浪费时间甚至过拟合。我一般会在脚本里加一个早停计数器连续 10 个 epoch 没有最佳 acc 更新就终止训练并自动恢复历史最佳权重。这样配合 checkpoint 机制训练过程非常稳妥集体请假时也敢放心挂着跑。6.3 多模型投票与导出部署数据量不大时单个模型的预测不够稳定我习惯训练三个不同随机种子的 B0 模型推理时对结果投票。这个改动对最终精度的提升很可观而且只需要在测试阶段写一段投票逻辑。导出部署也很直接把 model.eval() 之后的模型用 torch.jit.trace 打包成 TorchScript线上推理时不再依赖 Python 训练环境加载速度也会快很多。从那以后我每次跑迁移学习都会强制走一遍这个流程先确认数据目录无误再检查 num_classes 和权重路径最后启动训练盯住 val_acc 变化曲线。这套演示脚本帮我避开了最蠢的那些翻车希望帮到你。本文还有配套的精品资源点击获取
返回列表