ViT 微调指南用 4 个杠杆突破 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-models预训练权重一加载验证准确率就卡在 90% 上不去——这是 ViT 微调里最常见的困境模型本身没问题错的是微调的配方。本文基于 pytorch-image-modelstimm收录 ViT、ResNet、Swin、EfficientNet 等最大规模的 PyTorch 图像编码器/骨干网络附训练、评估、推理脚本与预训练权重梳理一套可复现的调法核心关键词只有一个——ViT 微调其余都是围绕它做参数取舍。微调 ViT 前先做的 3 个判断 ⚙️参数是最后才动的东西顺序错了怎么调都白调。先看三个前置决策。判断 1数据量决定解冻范围。ViT源码见 timm/models/vision_transformer.py由三部分组成Patch 嵌入层、多层注意力块构成的 Transformer 编码器、分类头。数据少几千张以内时冻结底层、只调顶层注意力块和分类头最稳数据量充足时直接全参微调上限更高。判断 2算力决定学习率与批次。批次大小决定学习率的安全上界——批次越大学习率越敢往上顶。GPU 吃紧时要么降批次配更低学习率要么开混合精度换吞吐别硬跑。判断 3目标决定正则强度。追求固定测试集上的绝对精度正则往重了配DropPath 到 0.2追求泛化与鲁棒性正则放轻、把力气花在数据增强上。一条端到端的最小训练链环境到调度五步代码只留关键参数完整参考 train.py。1. 环境git clone https://gitcode.com/GitHub_Trending/py/pytorch-image-models pip install -r pytorch-image-models/requirements.txt2. 数据timm/data/ 提供工厂函数224 是 ViT 的标准输入尺寸。from timm.data import create_dataset, create_loader dataset create_dataset(name, rootdata/train, splittrain) loader create_loader( dataset, input_size(3, 224, 224), batch_size32, is_trainingTrue, # 训练态下增强默认开启 )3. 模型换分类头 两项正则。import timm model timm.create_model( vit_base_patch16_224, pretrainedTrue, num_classes10, # 自定义类别数分类头随之重建 drop_rate0.1, # 全连接层 Dropout drop_path_rate0.1, # DropPath 随机深度建议区间 0.1–0.2 )4. 优化器AdamW 是微调的默认选择小学习率 明显权重衰减。from timm.optim import create_optimizer_v2 optimizer create_optimizer_v2( model, optadamw, lr5e-5, # 微调经验区间 5e-5 ~ 1e-4 weight_decay0.05, betas(0.9, 0.999), )5. 调度余弦退火 预热入口在 timm/scheduler/scheduler_factory.py。from timm.scheduler import create_scheduler_v2 scheduler, num_epochs create_scheduler_v2( optimizer, schedcosine, num_epochs30, warmup_epochs5, lr_min1e-6, # 余弦衰减的下限 warmup_lr1e-6, # 预热起点 )4 条准确率提升杠杆 按影响力从大到小排每条先说它解决什么再给取值与权衡。杠杆 1DropPath Dropout正则化编码器本体。解决训练涨、验证不涨。drop_path_rate控制在0.1–0.2它按深度线性递增地随机丢弃整个注意力块是 ViT 上最有效的正则drop_rate取0.1作用在分类头的全连接上。权衡超过 0.3 小数据集收敛会变慢先 0.1 起步观察到过拟合再加。杠杆 2余弦调度 预热。解决两头的问题——前期分类头刚重建、梯度剧烈预热阶段从warmup_lr线性爬升到峰值 lr 帮它稳住后期参数接近局部最优余弦曲线缓慢压低到lr_min给模型细调空间、有助于跳出局部最优。权衡预热占总轮数的 20%–40%即可30 轮给 5 轮预热刚好轮数少就别硬拉长。杠杆 3数据增强 随机擦除。解决模型背纹理。在 timm/data/transforms_factory.py 的create_transform或直接传给create_loader里配auto_augmentrand-m9-mstd0.5-inc1RandAugment多数据集验证有效、color_jitter0.4、re_prob0.25配re_modepixel做随机擦除、interpolationbicubic。权衡增强和 DropPath 都在正则两者同时拉满容易互相抵消先增强、后补 DropPath。杠杆 4模型 EMA 标签平滑。解决最后 1 个点被批次噪声吃掉和输出分布过度自信。EMA 实现见 timm/utils/model_ema.pyfrom timm.utils import ModelEmaV3 from timm.loss import LabelSmoothingCrossEntropy model_ema ModelEmaV3(model, decay0.9998, foreachTrue) criterion LabelSmoothingCrossEntropy(smoothing0.1)训练循环里每个 batch 反向传播、优化器步进之后顺手更新一次 EMAfor inputs, labels in loader: loss criterion(model(inputs), labels) loss.backward() optimizer.step() model_ema.update(model) # 每 batch 更新 EMAsmoothing0.1是安全默认值EMA 的decay别低于 0.999否则平均窗口太短、稳不住。怎么确认调对了基线先行先用冻结骨干、只训练分类头跑一版得到 linear probe 基线全参微调的每个杠杆都拿它做对照提升没有 1–2 个点就别算数。看 EMA 的 top-1不看训练 loss验证永远在model_ema上做with torch.no_grad()包起来过一遍验证集。早停验证精度连续5–10 轮不再刷新就停回滚历史最优权重别信再跑几轮会回来。曲线形态判据调对了的标志是训练/验证两条曲线贴得近且同步走训练明显甩开验证说明正则还不够回到杠杆 1 和杠杆 3。症状-原因-处方 排错速查 症状可能原因处方过拟合训练涨、验证掉正则不足 / 数据量小drop_path_rate加到0.2学习率降档或weight_decay加大增强策略加码早停 patience 取5–10训练不稳定loss 震荡、尖峰学习率过高、梯度爆炸学习率降到3e-5梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)检查归一化是否用了 ViT 默认的 ImageNet 均值/标准差推理慢FP32 全速跑、未编译混合精度推理torch.cuda.amp.autocast()model torch.compile(model)必要时降批次或换更小的模型变体下一步更大模型、混合精度、蒸馏换更大骨干vit_large_patch16_224数据够时上限更高但学习率要再压一档。混合精度训练autocastGradScaler开起来同卡吞吐通常翻倍省下的时间拿去多扫两组参数。知识蒸馏让大教师模型的软标签监督小 ViT小数据场景收益明显。版本timm 的模型库和调度器更新很快跑重要实验前先pip install --upgrade timmAPI 变动可以看 hfdocs/source/ 里的说明。杠杆的顺序就是投入的顺序正则、调度、增强、EMA 依次加码每加一层都用基线验证一次比一口气全开再猜哪层有问题要快得多。【免费下载链接】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),仅供参考