timm图像模型微调实战:小数据集上提升ViT精度的6个关键参数
【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 & V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models
预训练ViT权重下载到手,自己在数据集上精度却卡在90%上不去?pytorch-image-models(timm)收录了200多个图像backbone,并配套了完整的训练、评估、导出脚本。本文只讲6个关键参数怎么配,把小数据集上的精度从90%+推到97%左右。
项目速览:timm是PyTorch最大的图像backbone集合
timm的定位很明确:一个名字对应一个可用的图像编码器,且官方预训练权重直接可用。
- 覆盖面:ResNet、EfficientNet、ViT、Swin、ConvNeXt、MobileNetV3、RegNet等主流backbone都在 timm/models/ 下,用
list_models()就能查到全量清单 - 不止推理:仓库根目录的 train.py、
validate.py、inference.py、onnx_export.py构成完整的训练—评估—部署链路 - 适合谁:手头有自定义分类数据、想换backbone对比效果、需要导出ONNX上线的工程师
- 不适合谁:只需要单张模型文件做一次性推理的用户,直接
pip install timm即可,不需要整个仓库
快速跑起来:3步跑通预训练ViT推理
先确认环境和模型都正常,再谈调参。
pip install timm一行装好,权重在首次使用时自动下载,无需手动管理。
import torch import timm model = timm.create_model('vit_base_patch16_224', pretrained=True) model.eval() x = torch.randn(1, 3, 224, 224) with torch.no_grad(): logits = model(x) print(logits.shape) # torch.Size([1, 1000])这段代码解决“模型能不能跑”的问题:输入必须是3通道224×224,归一化用ImageNet均值方差。如果这步报显存或尺寸错误,先检查输入形状,别急着调参数。
核心机制拆解:三个工厂让微调变成流水线作业
timm把训练拆成三个工厂函数,官方脚本本质上就是组装它们,所以你要做的是换参数,而不是重写训练循环。
create_model:按名称查注册表,返回“架构+预训练配置”实例。ViT的实现就在 timm/models/vision_transformer.py,补丁嵌入、注意力块、分类头都在这一个文件里create_optimizer_v2:在标准优化器之上自动做两件事——BN和偏置参数豁免权重衰减、支持按层衰减学习率,这俩正是微调ViT最需要的create_scheduler_v2:cosine/step/poly等多种曲线,warmup内置,不用自己拼warmup逻辑
理解这一点后,后文所有调参都落在这三个函数的入参上,改哪一层、影响什么,一眼可见。
关键配置详解:最先要动的6个参数
微调时真正决定成败的是这6个参数,其余保持默认即可。
| 参数 | 设置位置 | 推荐起步值 | 为什么是这个值 |
|---|---|---|---|
num_classes | create_model | 你的实际类别数 | 决定分类头维度,不匹配直接报错,是第一个要改的 |
drop_path_rate | create_model | 0.1 | 随机深度是ViT微调的主力正则项;数据少于1万张用0.1,超过5万张降到0.05 |
drop_rate | create_model | 0.0 | 预训练权重已经稳定,再叠dropout收益小于代价 |
smoothing | LabelSmoothingCrossEntropy | 0.1 | 压低过度自信的概率,缓解小数据上过拟合 |
lr | create_optimizer_v2 | 3e-5(AdamW) | 微调学习率要比从头训练(约1e-3)低1个数量级,否则预训练特征会被冲掉 |
weight_decay | create_optimizer_v2 | 0.1 | 只作用于权重矩阵,BN/偏置自动豁免 |
模型侧配置:
model = timm.create_model( 'vit_base_patch16_224', pretrained=True, num_classes=24, # 换成你的类别数 drop_path_rate=0.1, # 小数据集主力正则项 )它解决“分类头维度”和“过拟合”两个问题。精度上不去时,优先动drop_path_rate,其次才是学习率。
优化器与调度器:
from timm.optim import create_optimizer_v2 from timm.scheduler import create_scheduler_v2 opt = create_optimizer_v2(model, opt='adamw', lr=3e-5, weight_decay=0.1) sched, epochs = create_scheduler_v2( opt, sched='cosine', num_epochs=20, warmup_epochs=3, warmup_lr=1e-6, min_lr=1e-7, )warmup前3个epoch学习率从1e-6爬升到3e-5,避免初期大梯度破坏预训练权重。训练前几个epoch掉精度就加warmup轮数,详见 timm/scheduler/scheduler_factory.py 的全部可选项。
损失函数:
from timm.loss import LabelSmoothingCrossEntropy loss_fn = LabelSmoothingCrossEntropy(smoothing=0.1)一行替换标准交叉熵即可。验证精度忽高忽低时,先确认训练和验证用的是不是同一个loss口径。
进阶调优:再抠出1~2个百分点的3个技巧
按层衰减学习率:让浅层学得更慢
ViT的微调讲究“浅层保留、深层适应”。create_optimizer_v2原生支持层衰减:
opt = create_optimizer_v2(model, opt='adamw', lr=1e-4, weight_decay=0.1, layer_decay=0.75)每往浅层走一层,学习率乘以0.75,最底层约为顶层的1/10。当发现特征层可视化几乎不变、只有头在动时,把layer_decay提到0.8~0.9。
EMA warmup:短训练别用固定decay
ModelEmaV3的默认decay=0.9999是为长训练设计的,短训练下EMA会被前期权重拖慢。
from timm.utils import ModelEmaV3 ema = ModelEmaV3(model, decay=0.9999, use_warmup=True, device='cpu') # 每个优化器step之后调用: ema.update(model)use_warmup=True让decay随步数逐渐爬升,训练不足5万步时务必开启;device='cpu'省一半显存,代价是每步一次PCIe传输。实现细节在 timm/utils/model_ema.py。
梯度裁剪:发散的最后一道保险
微调初期若loss出现尖峰,说明某一步梯度过大。参照 train.py 的做法,在反传后加:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)从1.0起步。loss尖峰频率下降后,可逐步放宽到3.0,不必长期收紧。
踩坑与解法:5个高频问题
| 现象 | 原因与解法 |
|---|---|
| 前1~2个epoch精度先跌后涨 | 正常,AdamW+cosine需要warmup;把warmup_epochs提到3~5 |
| ⚠️ EMA模型验证精度反而更低 | decay起步太猛,加use_warmup=True;仍不行就改用原模型验证 |
| 加EMA后显存直接翻倍告警 | device='cpu'把EMA挪到内存,接受少量耗时换显存 |
| 旧版checkpoint加载报 unexpected keys | 用strict=False加载,或核对timm版本与权重版本是否匹配 |
| 老卡上fp16训练loss变NaN | Ampere及以上直接上bf16,省掉loss scaler;老卡把clip_grad收紧到1.0 |
以上问题里,80%的精度异常出在前两条,动手改模型结构前先查warmup和EMA配置。
收尾
先用create_model确认预训练权重能正常出结果,再把6个参数按表格落到create_model、create_optimizer_v2、LabelSmoothingCrossEntropy三个入口,最后用EMA和层衰减抠尾点。整套流程不用写训练框架,改参数即可复现。
延伸方向:
- 输入分辨率从224提到384,配合
img_size=(3, 384, 384)和位置插值,精度普遍再涨1~2点 - 同一数据上横评ConvNeXt、Swin、ViT三类backbone,用
list_models()挑参数量相近的做公平对比 - 用仓库里的
onnx_export.py把调好的模型导出ONNX,再跑onnx_validate.py对齐精度
【免费下载链接】pytorch-image-modelsThe largest collection of PyTorch image encoders / backbones. Including train, eval, inference, export scripts, and pretrained weights -- ResNet, ResNeXT, EfficientNet, NFNet, Vision Transformer (ViT), MobileNetV4, MobileNet-V3 & V2, RegNet, DPN, CSPNet, Swin Transformer, MaxViT, CoAtNet, ConvNeXt, and more项目地址: https://gitcode.com/GitHub_Trending/py/pytorch-image-models
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考