尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

timm图像模型微调实战:小数据集上提升ViT精度的6个关键参数

timm图像模型微调实战:小数据集上提升ViT精度的6个关键参数 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-modelstimm收录了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, pretrainedTrue) 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_v2cosine/step/poly等多种曲线warmup内置不用自己拼warmup逻辑理解这一点后后文所有调参都落在这三个函数的入参上改哪一层、影响什么一眼可见。关键配置详解最先要动的6个参数微调时真正决定成败的是这6个参数其余保持默认即可。参数设置位置推荐起步值为什么是这个值num_classescreate_model你的实际类别数决定分类头维度不匹配直接报错是第一个要改的drop_path_ratecreate_model0.1随机深度是ViT微调的主力正则项数据少于1万张用0.1超过5万张降到0.05drop_ratecreate_model0.0预训练权重已经稳定再叠dropout收益小于代价smoothingLabelSmoothingCrossEntropy0.1压低过度自信的概率缓解小数据上过拟合lrcreate_optimizer_v23e-5AdamW微调学习率要比从头训练约1e-3低1个数量级否则预训练特征会被冲掉weight_decaycreate_optimizer_v20.1只作用于权重矩阵BN/偏置自动豁免模型侧配置model timm.create_model( vit_base_patch16_224, pretrainedTrue, num_classes24, # 换成你的类别数 drop_path_rate0.1, # 小数据集主力正则项 )它解决“分类头维度”和“过拟合”两个问题。精度上不去时优先动drop_path_rate其次才是学习率。优化器与调度器from timm.optim import create_optimizer_v2 from timm.scheduler import create_scheduler_v2 opt create_optimizer_v2(model, optadamw, lr3e-5, weight_decay0.1) sched, epochs create_scheduler_v2( opt, schedcosine, num_epochs20, warmup_epochs3, warmup_lr1e-6, min_lr1e-7, )warmup前3个epoch学习率从1e-6爬升到3e-5避免初期大梯度破坏预训练权重。训练前几个epoch掉精度就加warmup轮数详见 timm/scheduler/scheduler_factory.py 的全部可选项。损失函数from timm.loss import LabelSmoothingCrossEntropy loss_fn LabelSmoothingCrossEntropy(smoothing0.1)一行替换标准交叉熵即可。验证精度忽高忽低时先确认训练和验证用的是不是同一个loss口径。进阶调优再抠出1~2个百分点的3个技巧按层衰减学习率让浅层学得更慢ViT的微调讲究“浅层保留、深层适应”。create_optimizer_v2原生支持层衰减opt create_optimizer_v2(model, optadamw, lr1e-4, weight_decay0.1, layer_decay0.75)每往浅层走一层学习率乘以0.75最底层约为顶层的1/10。当发现特征层可视化几乎不变、只有头在动时把layer_decay提到0.8~0.9。EMA warmup短训练别用固定decayModelEmaV3的默认decay0.9999是为长训练设计的短训练下EMA会被前期权重拖慢。from timm.utils import ModelEmaV3 ema ModelEmaV3(model, decay0.9999, use_warmupTrue, devicecpu) # 每个优化器step之后调用 ema.update(model)use_warmupTrue让decay随步数逐渐爬升训练不足5万步时务必开启devicecpu省一半显存代价是每步一次PCIe传输。实现细节在 timm/utils/model_ema.py。梯度裁剪发散的最后一道保险微调初期若loss出现尖峰说明某一步梯度过大。参照 train.py 的做法在反传后加torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)从1.0起步。loss尖峰频率下降后可逐步放宽到3.0不必长期收紧。踩坑与解法5个高频问题现象原因与解法前1~2个epoch精度先跌后涨正常AdamWcosine需要warmup把warmup_epochs提到3~5⚠️ EMA模型验证精度反而更低decay起步太猛加use_warmupTrue仍不行就改用原模型验证加EMA后显存直接翻倍告警devicecpu把EMA挪到内存接受少量耗时换显存旧版checkpoint加载报 unexpected keys用strictFalse加载或核对timm版本与权重版本是否匹配老卡上fp16训练loss变NaNAmpere及以上直接上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),仅供参考
返回列表