ARTICLE DETAIL

资讯详情

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

ViT 微调实战指南:timm 的三组旋钮与调参速查,让自定义数据集的准确率不再趴着不动

ViT 微调实战指南:timm 的三组旋钮与调参速查,让自定义数据集的准确率不再趴着不动 ViT 微调实战指南timm 的三组旋钮与调参速查让自定义数据集的准确率不再趴着不动【免费下载链接】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-modelstimmpytorch-image-models是目前 PyTorch 生态里最大的图像主干网络集合ViT、Swin、ConvNeXt、EfficientNet 等上百种模型都带预训练权重和开箱即用的训练/验证脚本。拿到一个 ImageNet 预训练好的 ViT想在几千张自定义数据上微调出好效果卡点往往不在模型本身而在几个参数的搭配。这篇按「先把流程跑通 → 再逐个拧旋钮 → 最后查表排错」的顺序讲一遍完整做法。微调本质三组旋钮不管数据集多特殊ViT 微调基本就是在拧三组东西优化器旋钮用什么优化器、学习率多大、权重衰减多强——决定走得快不快、稳不稳。数据旋钮增强策略、随机擦除、插值方式——决定模型看到的监督信号有多少有效信息。正则化旋钮DropPath、标签平滑、模型 EMA——决定泛化掉不掉链子。先建立这个全局认知后面每个参数你都能归位它到底是在解决收敛问题还是过拟合问题。ViT 的结构实现在 timm/models/vision_transformer.py图像先切成 patch 序列过 Transformer 编码器最后进分类头。微调时数据量够比如上万张就全量微调只有几百张时可以考虑只放开分类头和最后几个 block学习率再压低一档。先跑通一遍最小可运行示例环境一行搞定pip install timm下面这段把模型、数据、优化器、调度器、EMA 全部串起来可以直接当模板import timm import torch from timm.data import create_dataset, create_loader train_loader create_loader( create_dataset(name, rootdata, splittrain, is_trainingTrue), input_size(3, 224, 224), batch_size64, is_trainingTrue, auto_augmentrand-m9-mstd0.5-inc1, re_prob0.25, re_modepixel, interpolationbicubic) model timm.create_model(vit_base_patch16_224, pretrainedTrue, num_classes10, drop_path_rate0.1, drop_rate0.1) optimizer timm.optim.create_optimizer_v2(model, optadamw, lr1e-4, weight_decay0.05) scheduler, _ timm.scheduler.create_scheduler_v2( optimizer, schedcosine, num_epochs30, warmup_epochs5, min_lr1e-6) ema timm.utils.ModelEmaV3(model)训练主循环不到 10 行criterion timm.loss.LabelSmoothingCrossEntropy(smoothing0.1) for epoch in range(30): model.train() for x, y in train_loader: x, y x.cuda(), y.cuda() loss criterion(model(x), y) optimizer.zero_grad() loss.backward() optimizer.step() ema.update(model) scheduler.step() # cosine 默认按 epoch 步进验证时用ema.module而不是modelEMA 权重才是你最终该保存的那份。优化器怎么选AdamW 的两个关键值创建入口是create_optimizer_v2timm/optim/_optim_factory.py。选optadamw两个值重点盯学习率lrViT 微调建议从5e-5 ~ 2e-4起。数据少、batch 小就取小值batch 大≥128可以放到 2e-4。权重衰减weight_decay0.05对 Transformer 类模型这是标准值配合 AdamW 的解耦衰减能有效压住过拟合。一个原文档里容易忽略的细节filter_bias_and_bnTrue是默认值框架会自动给 bias 和归一化层参数豁免权重衰减所以你不需要手动分参数组直接传一个weight_decay就行。学习率调度余弦退火怎么配调度器工厂在 timm/scheduler/scheduler_factory.py微调的默认推荐是schedcosine理由前期学习率高、快速适配新数据后期平滑降到低位、精修决策边界不需要你手动找衰减点。warmup_epochs3~5预热防止前几个 epoch 把预训练权重冲坏数据越小越需要。min_lr1e-6余弦退到的地板值一般取峰值学习率的 1%~10%。step_on_epochs默认True每个 epoch 末调一次scheduler.step()即可如果按 step 更新需要额外传updates_per_epoch。数据增强RandAugment 随机擦除怎么给增强配置集中在create_loader/create_transformtimm/data/transforms_factory.py两个参数值得单独调auto_augmentrand-m9-mstd0.5-inc1m9 个增强操作、mstd0.5 强度是 ImageNet 训练验证过的默认配方。数据少可以降到rand-m6-mstd0.5-inc1。随机擦除三件套参数推荐值作用re_prob0.2525% 的图会被挖一块逼模型别只依赖局部纹理re_modepixel用像素均值填充比纯黑块更自然re_count1每张图挖 1 块即可类别少时可升到 2另外两个容易踩的默认值interpolation训练侧给bicubicViT 对分辨率敏感双线性会丢细节归一化的 mean/std 默认就是 ImageNet 的 0.485/0.229不要换成自己数据的均值否则预训练权重直接作废。正则化三件套DropPath、标签平滑、EMADropPath随机深度建模型时传drop_path_rate0.1训练时每个 block 按线性递增的概率被整体跳过。过拟合严重就加到0.2。分类头的drop_rate0.1可以顺手带上。标签平滑timm/loss/cross_entropy.py 里的LabelSmoothingCrossEntropysmoothing0.1。把全押一类的硬标签变软模型不敢过度自信验证精度通常能白捡 0.5%~1%。模型 EMAtimm/utils/model_ema.py 的ModelEmaV3每步维护一份权重的指数滑动平均天然滤掉单次梯度的抖动ema timm.utils.ModelEmaV3(model) # decay 默认 0.9999 # 每个 batch 优化器 step 之后 ema.update(model) # 验证 / 导出时用 ema.module短训练几十个 epoch可以把decay调低到0.999~0.9998让 EMA 更快跟上当前权重显存紧张时传devicecpu把平均权重放 CPU 上。调参速查表参数推荐起点过拟合时欠拟合时lrAdamW1e-4降到 5e-5升到 2e-4weight_decay0.05升到 0.1降到 0.01drop_path_rate0.1升到 0.2降到 0.05smoothing0.1升到 0.15降到 0.05EMAdecay0.9998维持 0.9999降到 0.999warmup_epochs5升到 83 或去掉min_lr1e-6维持升到 1e-5re_prob0.25升到 0.4降到 0.15auto_augmentrand-m9-mstd0.5-inc1维持rand-m6-mstd0.5-inc1原则每次只动一个旋钮跑 3~5 个 epoch 看趋势再决定方向别一次改一堆然后不知道是谁的功劳。排错手册症状大概率原因对策loss 直接 nan / 爆炸学习率过高或没预热lr 减半warmup 加到 5~8 个 epoch仍不行就torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)训练 acc 99%验证 acc 上不去正则化不足在背数据drop_path_rate0.2 增强档位上调数据本身少就先别再硬堆了loss 几乎不动num_classes和数据标签数不符、归一化被改过、参数被冻结核对类别映射文件确认 mean/std 是 ImageNet 值print([p.requires_grad for p in model.parameters()])抽查指标来回抖、收敛路径锯齿状单份权重噪声大验证统一切到ema.module通常立刻平稳验证/推理太慢全精度跑 无编译推理套torch.amp.autocast(cuda)model torch.compile(model)再试进阶方向分层权重衰减 / 分层学习率create_optimizer_v2支持layer_decay参数给浅层更小的学习率深层更充分地适配对深层 ViT 收益明显。知识蒸馏timm 自带 timm/task/ 下的蒸馏模块distillation.py、token_distillation.py让小模型跟着大老师走比继续加数据便宜得多。更大分辨率vit_base_patch16_224换成vit_base_patch16_384配合把input_size提到 384细粒度分类任务常有惊喜代价是显存翻倍。【免费下载链接】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),仅供参考
返回列表