ARTICLE DETAIL

资讯详情

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

ViT 微调从 90% 推到 98%:用 timm 三轮调参拉满预训练 ViT

ViT 微调从 90% 推到 98%:用 timm 三轮调参拉满预训练 ViT ViT 微调从 90% 推到 98%用 timm 三轮调参拉满预训练 ViT【免费下载链接】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用 timmpytorch-image-models对预训练 ViT 做微调时第一遍训练很难直接达标。下面按三个真实症状——验证曲线震荡、小数据集过拟合、收敛慢——给出学习率、数据增强和 EMA 的对应调法参数值可直接照抄。微调前的三个症状先定性再动手先说结论改参数之前先判断你属于哪一种症状否则调错方向只会浪费训练轮数。症状 A验证 loss 前期锯齿状震荡甚至先升后降。十有八九是学习率偏大或没有预热模型在破坏预训练权重。症状 B训练 loss 快速贴地验证 loss 在 2~3 个 epoch 后掉头向上。典型的过拟合数据集规模撑不起当前模型的表达能力。症状 C训练和验证都缓慢下降30 个 epoch 还到不了预训练模型的水平。通常是学习率过小、调度太保守或者数据管道里增强被意外关闭。ViT 的结构可以简单理解为「patch 嵌入 → 多层自注意力编码器 → 分类头」源码在 timm/models/vision_transformer.py。微调的本质是数据量小就少动底层权重数据量大就放开让所有层重新适应。ViT 微调学习率怎么设AdamW 加余弦预热是稳妥组合先给结论AdamW、峰值学习率 5e-5 到 1e-4、权重衰减 0.05、余弦衰减配 3~5 个 epoch 的 warmup是自定义数据集上的安全起点。这个区间下症状 A 基本可以根除。加载预训练权重时顺手把随机深度stochastic depth打开它是最便宜的正则化手段import timm model timm.create_model( vit_base_patch16_224, pretrainedTrue, num_classes10, drop_rate0.1, drop_path_rate0.1, )drop_rate作用于分类头前的 dropoutdrop_path_rate控制按层跳过的概率。两者从 0.1 起步过拟合严重时再往 0.2 加。优化器和调度器都用 timm 的工厂函数参数组会自动把 bias 和 norm 层排除在权重衰减之外from timm.optim import create_optimizer_v2 optimizer create_optimizer_v2( model, optadamw, lr8e-5, weight_decay0.05, )学习率取区间中值 8e-5数据特别少时压到 5e-5数据量大时可放到 1e-4。from timm.scheduler import create_scheduler_v2 scheduler, num_epochs create_scheduler_v2( optimizer, schedcosine, num_epochs30, warmup_epochs5, min_lr1e-6, warmup_lr1e-6, )warmup 前 5 个 epoch 把学习率从 1e-6 线性拉到峰值之后余弦衰减到 1e-6。预热阶段是治症状 A 的关键——没有它前几个 batch 的大步长会直接冲坏预训练特征。调度器参数定义见 timm/scheduler/scheduler_factory.py。timm 数据增强配置三个开关直接决定泛化上限结论训练集用 RandAugment、随机擦除概率 0.25、bicubic 插值验证集只做 resize 和归一化其余全关。增强只在训练路径生效验证路径加了增强你测出来的精度会系统性偏低先排除这个低级错误。create_transform一条调用就能同时覆盖训练和验证两套管线from timm.data import create_transform train_tf create_transform( input_size(3, 224, 224), is_trainingTrue, auto_augmentrand-m9-mstd0.5-inc1, color_jitter0.4, re_prob0.25, re_modepixel, re_count1, interpolationbicubic, ) val_tf create_transform(input_size(3, 224, 224), is_trainingFalse)auto_augmentrand-m9-mstd0.5-inc1是 RandAugment 的 m9 强度配置在 ImageNet 系模型上验证充分re_prob0.25的像素级随机擦除强迫 ViT 不能依赖单一区域做判断对 patch 类模型尤其有效。完整参数在 timm/data/transforms_factory.py 里都能查到插值务必用 bicubic——ViT 预训练时就是这么 resize 的验证集用 bilinear 会在小目标上白白掉零点几个点。数据集侧用create_dataset加create_loader组装即可自定义类别放一个 class_map 文本文件避免依赖目录命名。预训练模型过拟合怎么办EMA 加标签平滑双保险结论症状 B 的解法不是减数据而是「EMA 标签平滑 更大的 drop_path」这套组合拳。单独上任何一项收益都有限叠起来验证精度通常能再抬 1~2 个点。EMA 是训练过程中维护一份参数的指数滑动平均验证时用它而不是原始权重能显著抹平单 batch 噪声带来的抖动。timm 里直接用ModelEmaV3from timm.utils import ModelEmaV3 model_ema ModelEmaV3(model, decay0.9998, devicecuda) # 训练循环内每个 batch 之后 model_ema.update(model)decay0.9998适合几千到几万步规模的微调若训练步数很少几千步以内可以降到 0.999。注意验证和最终导出权重都要走model_ema.module这是最容易漏的一步源码在 timm/utils/model_ema.py。标签平滑防止 ViT 对预训练见过的类别过度自信小数据集上收益明显from timm.loss import LabelSmoothingCrossEntropy criterion LabelSmoothingCrossEntropy(smoothing0.1)平滑系数 0.1 是社区默认值类别极少少于 20 类时可提到 0.15。决策表什么场景配什么值把上面三轮调参浓缩成一张表照着改就行场景峰值学习率drop_path_rate增强策略训练轮数小数据5k 张5e-50.1~0.2RandAugment 随机擦除 0.2540~60warmup 占 1/6中数据5k~50k8e-50.1同上30warmup 5 epoch大数据50k1e-40.0~0.1同上前两项即可30~50调度放开症状 C 收敛慢上调一档学习率先排除增强被关闭不动检查is_trainingTrue不动调参顺序上第一轮只动学习率和 warmup第二轮动正则化drop_path、平滑第三轮才动增强强度每轮记录验证曲线再决定下一步。一轮改三处出了问题无法归因。下一步可以试什么症状都解决之后还有三件事值得排队混合精度训练bf16 在 A100 及以上卡上对 ViT 提速明显且几乎无损配合torch.cuda.amp即可。换更大的模型变体vit_large_patch16_224在数据充足时上限更高但学习率要按上面决策表往小数据档靠。梯度裁剪兜底若 warmup 后仍偶发尖峰加一行torch.nn.utils.clip_grad_norm_max_norm 取 1.0成本几乎为零。完整的端到端脚本可以直接参考仓库根目录的train.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),仅供参考
返回列表