ARTICLE DETAIL

资讯详情

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

MAE 图像分类微调实战指南:从预训练权重评估、端到端 Fine-tuning 到 Linear Probing(PyTorch 实现)

MAE 图像分类微调实战指南:从预训练权重评估、端到端 Fine-tuning 到 Linear Probing(PyTorch 实现) 人工智能计算机视觉深度学习机器学习预训练微调【免费下载链接】maePyTorch implementation of MAE https//arxiv.org/abs/2111.06377项目地址https://gitcode.com/gh_mirrors/ma/mae点击查看免费下载本文基于本仓库Masked Autoencoders: A PyTorch ImplementationMAE 论文《Masked Autoencoders Are Scalable Vision Learners》的 PyTorch/GPU 复现实现的 FINETUNE.md 整理而成并结合 main_finetune.py、main_linprobe.py、models_vit.py、engine_finetune.py 等源码对命令背后的实现原理做纵深剖析。读完本文你将掌握如何使用官方微调权重快速评估模型并复现 ImageNet 精度如何用 4 节点 / 单节点分布式训练对 MAE 预训练模型做端到端分类微调以及如何用 Linear Probing 在冻结主干的情况下评估预训练特征质量。所有命令均以本仓库代码为准可直接复制运行。一、前置准备1.1 仓库结构概览与微调相关的核心文件如下均位于仓库根目录或util/下FINETUNE.md本文档微调、评估与线性探测的官方操作说明main_finetune.py端到端微调 / 评估主程序含全部命令行参数定义main_linprobe.py线性探测主程序engine_finetune.py训练一个 epoch 与评估Acc1/Acc5的实现models_vit.py支持 global pooling 的 ViT 模型定义Base/Large/Hugeutil/lr_decay.pyBEiT 风格的层间学习率衰减layer-wise lr decayutil/lr_sched.pywarmup 半周期余弦学习率调度util/pos_embed.py位置编码加载与插值适配不同输入分辨率util/datasets.pyImageNet 数据集与训练/评估数据增强管线submitit_finetune.py、submitit_linprobe.py基于 submitit 的 Slurm 多节点作业提交脚本1.2 环境与依赖PyTorch CUDA GPU本仓库基于 PyTorch GPU 复现原版实现为 TensorFlow TPUtimm0.3.2main_finetune.py第 26 行有硬性版本断言assert timm.__version__ 0.3.2请务必锁定该版本torchvision、tensorboard多节点训练需要额外安装 submititpip install submitit单节点训练则不需要1.3 数据集目录结构${IMAGENET_DIR}是一个包含{train, val}两个子目录的 ImageNet 目录${IMAGENET_DIR}/ ├── train/ # 训练集按类别分目录torchvision ImageFolder 格式 └── val/ # 验证集在 util/datasets.py 中build_dataset通过datasets.ImageFolder(root, transformtransform)加载数据root 分别为${IMAGENET_DIR}/train与${IMAGENET_DIR}/val。二、评估官方微调权重Evaluation作为微调正确性的 sanity check首先用官方公开的 ImageNetfine-tuned权重做纯评估。下表列出了官方三个规模微调后的权重、md5 与参考精度数据来自 FINETUNE.md项目ViT-BaseViT-LargeViT-Hugefine-tuned checkpointmae_finetuned_vit_base.pthmae_finetuned_vit_large.pthmae_finetuned_vit_huge.pthmd51b25e951f5502541f2reference ImageNet accuracy83.66485.95286.928权重文件托管于官方公开下载地址dl.fbaipublicfiles.com/mae/finetune/目录文件名与上表一致下载后请用 md5 校验文件完整性。2.1 单 GPU 评估 ViT-Basepython main_finetune.py --eval --resume mae_finetuned_vit_base.pth --model vit_base_patch16 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出* Acc1 83.664 Acc5 96.530 loss 0.7312.2 评估 ViT-Largepython main_finetune.py --eval --resume mae_finetuned_vit_large.pth --model vit_large_patch16 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出* Acc1 85.952 Acc5 97.570 loss 0.6462.3 评估 ViT-Hugepython main_finetune.py --eval --resume mae_finetuned_vit_huge.pth --model vit_huge_patch14 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出* Acc1 86.928 Acc5 98.088 loss 0.5842.4 评估流程的源码解析从 main_finetune.py 看评估路径的核心逻辑如下参数入口--eval开关第 137-138 行触发仅评估模式--resume指定权重路径第 132-133 行--model选择模型结构第 51-52 行默认vit_large_patch16--batch_size默认 64第 44-45 行官方示例统一用 16。模型构建models_vit.__dict__args.model第 227-231 行。--nb_classes默认 1000适配 ImageNet-1K。权重加载与结构适配第 233-257 行预训练权重中的head.weight/head.bias因分类头尺寸不同会被删除随后interpolate_pos_embed见 util/pos_embed.py在输入分辨率变化时用 bicubic 插值调整位置编码最后用trunc_normal_(model.head.weight, std2e-5)重新初始化分类头。由于仓库默认--global_pool第 114-115 行加载时的缺失键须恰好为head与fc_norm四组参数第 251-254 行的断言这从源码层面印证了 MAE 微调统一使用全局平均池化替代 class token。精度统计evaluate在 engine_finetune.py 中实现使用timm.utils.accuracy计算 Top-1 / Top-5并在torch.cuda.amp.autocast()下推理自动混合精度与训练保持一致。三、端到端微调Fine-tuning微调的输入是 MAE 自监督预训练权重官方预训练权重见 README.md 中的 pre-trained checkpoints 表格包括 ViT-Basemae_pretrain_vit_base.pth、ViT-Large、ViT-Hugemd5 分别为8cad7c、b8b06e、9bdbb0。官方预训练权重使用归一化像素损失--norm_pix_loss训练 1600 epoch论文 Table 3因此微调超参数与默认未归一化基线略有差异。3.1 多节点分布式训练4 节点 × 8 GPU使用 submitit 在 Slurm 集群上提交作业需先pip install submititpython submitit_finetune.py \ --job_dir ${JOB_DIR} \ --nodes 4 \ --batch_size 32 \ --model vit_base_patch16 \ --finetune ${PRETRAIN_CHKPT} \ --epochs 100 \ --blr 5e-4 --layer_decay 0.65 \ --weight_decay 0.05 --drop_path 0.1 --reprob 0.25 --mixup 0.8 --cutmix 1.0 \ --dist_eval --data_path ${IMAGENET_DIR}关键说明有效 batch sizebatch_size每 GPU 32×nodes4× 每节点 GPU 数81024。blr是基准学习率实际lr按线性缩放规则计算lr blr × 有效batch size / 256即5e-4 × 1024 / 256 2e-3。官方用 4 个不同随机种子跑了 4 次实验结果为 83.63、83.66、83.52、83.46均值 83.57标准差 0.08。训练时间约7 小时 11 分32 张 V100 GPU。3.2 ViT-Large 微调脚本python submitit_finetune.py \ --job_dir ${JOB_DIR} \ --nodes 4 --use_volta32 \ --batch_size 32 \ --model vit_large_patch16 \ --finetune ${PRETRAIN_CHKPT} \ --epochs 50 \ --blr 1e-3 --layer_decay 0.75 \ --weight_decay 0.05 --drop_path 0.2 --reprob 0.25 --mixup 0.8 --cutmix 1.0 \ --dist_eval --data_path ${IMAGENET_DIR}4 个随机种子结果为 85.95、85.87、85.76、85.88均值 85.87标准差 0.07。训练时间约8 小时 52 分32 张 V100。--use_volta32表示申请 32G 显存的 V100。3.3 ViT-Huge 微调脚本python submitit_finetune.py \ --job_dir ${JOB_DIR} \ --nodes 8 --use_volta32 \ --batch_size 16 \ --model vit_huge_patch14 \ --finetune ${PRETRAIN_CHKPT} \ --epochs 50 \ --blr 1e-3 --layer_decay 0.75 \ --weight_decay 0.05 --drop_path 0.3 --reprob 0.25 --mixup 0.8 --cutmix 1.0 \ --dist_eval --data_path ${IMAGENET_DIR}训练时间约13 小时 9 分64 张 V100。ViT-Huge 使用 patch14输入 224×224 时对应 16×16256 个 patchdrop_path提高到 0.3 以增强正则。3.4 单节点训练1 节点 × 8 GPU无需 submitit用torch.distributed.launch代替 submititOMP_NUM_THREADS1 python -m torch.distributed.launch --nproc_per_node8 main_finetune.py \ --accum_iter 4 \ --batch_size 32 \ --model vit_base_patch16 \ --finetune ${PRETRAIN_CHKPT} \ --epochs 100 \ --blr 5e-4 --layer_decay 0.65 \ --weight_decay 0.05 --drop_path 0.1 --mixup 0.8 --cutmix 1.0 --reprob 0.25 \ --dist_eval --data_path ${IMAGENET_DIR}关键说明有效 batch size 32每 GPU×accum_iter4× 8GPU1024。--accum_iter 4用梯度累积模拟 4 个节点在单机显存受限时也能凑出同样的有效 batch size。从 engine_finetune.py 可以看到梯度累积的具体实现loss / accum_iter且仅在(data_iter_step 1) % accum_iter 0时才通过loss_scaler(...)真正更新梯度并optimizer.zero_grad()。注意此处没传--use_volta32因为那是 submitit 专属参数在 submitit_finetune.py 中定义。3.5 核心参数速查表下表汇总 main_finetune.py 中与微调效果直接相关的关键参数及其默认值、含义参数默认值说明--modelvit_large_patch16模型结构vit_base_patch16/vit_large_patch16/vit_huge_patch14--batch_size64每 GPU batch size有效 batch batch_size × accum_iter × GPU 数--accum_iter1梯度累积迭代数用于在显存受限时增大有效 batch size--epochs50训练轮数Base 用 100Large/Huge 用 50--blr1e-3基准学习率实际 lr blr × 有效 batch / 256--layer_decay0.75层间学习率衰减系数Base 0.65Large/Huge 0.75--weight_decay0.05权重衰减--drop_path0.1DropPath 随机深度丢弃率Base 0.1Large 0.2Huge 0.3--reprob0.25Random Erasing 概率此处跟随 DeiT 设置--mixup0mixup alpha0 时启用--cutmix0cutmix alpha0 时启用--smoothing0.1标签平滑系数启用 mixup 时由 mixup 标签变换接管--aarand-m9-mstd0.5-inc1AutoAugment 策略--global_pool/--cls_tokenglobal_pool 默认开分类特征来源全局平均池化 / class token--finetune空预训练权重路径--data_path/datasets01/imagenet_full_size/061417/数据集根目录含 train/val--nb_classes1000分类类别数--warmup_epochs5学习率 warmup 轮数--min_lr1e-6余弦调度下界--clip_gradNone梯度裁剪范数默认不裁剪--dist_evalFalse分布式评估训练时推荐开启以加速监控--evalFalse仅评估模式--output_dir/--log_dir./output_dir模型保存目录 / TensorBoard 日志目录3.6 微调背后的源码级原理1层间学习率衰减layer-wise lr decay微调的核心技巧之一。在 main_finetune.py 中优化器参数组由 util/lr_decay.py 的param_groups_lrd构建越靠近输入层的参数 lr 越小按layer_decay^(num_layers - layer_id)缩放靠近分类头的层 lr 越大。层 id 由get_layer_id_for_vit第 64-77 行分配cls_token/pos_embed/patch_embed属于第 0 层blocks.i属于第 i1 层fc_norm与head属于最后一层。同时 1 维参数如 norm 的 weight/bias默认不做权重衰减。2学习率调度util/lr_sched.py 实现 warmup 半周期余弦衰减epoch warmup_epochs时线性升温之后按余弦曲线从lr降到min_lr并且每个参数组乘上各自的lr_scale。在 engine_finetune.py 中该调度是逐 iteration更新的而非逐 epoch保证不同 batch size 下的曲线可对齐。3数据增强util/datasets.py 的训练增强基于 timm 的create_transformAutoAugment默认rand-m9-mstd0.5-inc1、bicubic 插值、Random Erasing--reprob 0.25评估增强为 resize center crop。mixup / cutmix 在 main_finetune.py 中根据--mixup/--cutmix构建 timm 的Mixup对象配合SoftTargetCrossEntropy损失第 290-296 行启用 mixup 时用软标签交叉熵否则用LabelSmoothingCrossEntropy或普通 CE。4自动混合精度AMP整个前向与损失计算都在torch.cuda.amp.autocast()中进行梯度缩放使用仓库自实现的NativeScalerWithGradNormCountutil/misc.py。这是本 PyTorch/GPU 复现与原 TensorFlow/TPU 实现的重要系统差异之一详见 3.7 Notes。5提交脚本的工程细节submitit_finetune.py 通过submitit.AutoExecutor提交 Slurm 作业每节点申请 8 GPU、每 GPU 一个 task、每 task 10 CPU、内存 40GB/节点 GPU作业目录支持%j通配符自动按 job_id 落盘Trainer.checkpoint支持断点续跑检测到checkpoint.pth时自动--resume重排队。这些参数面向 Slurm 集群在非 Slurm 环境请使用 3.4 的单节点方式。3.7 微调注意事项Notes来自原文档归一化像素官方提供的预训练权重是用归一化像素损失--norm_pix_loss训练 1600 epoch 得到的论文 Table 3微调超参数因此与使用未归一化像素的默认基线略有不同。AMP 与数值行为差异原版 MAE 是 TensorFlowTPU 且无显式混合精度本复现为 PyTorchGPU 并启用 AMP两个平台存在数值行为差异。本仓库微调统一使用--global_pool全局平均池化用--cls_token效果相当但在 GPU 上微调 ViT-Huge 时有产生 NaN 的可能TPU 上未观察到。关闭 AMP 可以规避该问题但训练更慢。RandErase这里跟随 DeiT 设置--reprob 0.25其效果小于随机方差即对最终精度的贡献不显著。四、线性探测Linear ProbingLinear Probing 用于评估预训练表征质量冻结主干所有参数只训练一个线性分类头看特征本身的线性可分性。4.1 4 节点 × 8 GPU 训练 ViT-Basepython submitit_linprobe.py \ --job_dir ${JOB_DIR} \ --nodes 4 \ --batch_size 512 \ --model vit_base_patch16 --cls_token \ --finetune ${PRETRAIN_CHKPT} \ --epochs 90 \ --blr 0.1 \ --weight_decay 0.0 \ --dist_eval --data_path ${IMAGENET_DIR}关键说明有效 batch size 512 × 4 × 8 16384线性探测可以吃下超大 batch。blr基准学习率 0.1实际 lr 0.1 × 16384 / 256 6.4。--weight_decay 0.0源码注释说明following MoCo v1线性探测不使用权重衰减。--cls_token线性探测使用 class token 特征与微调的 global_pool 相反。训练时间约2 小时 20 分90 epoch32 张 V100。单节点训练方式同微调用torch.distributed.launch运行 main_linprobe.py并可配合--accum_iter累积梯度。4.2 训练 ViT-Large / ViT-Huge将--model改为vit_large_patch16或vit_huge_patch14--epochs 50即可大模型 50 epoch 已足够。4.3 线性探测的源码级实现与端到端微调相比main_linprobe.py 有四处关键差异弱增强第 131-141 行训练只用 RandomResizedCrop(224) 随机水平翻转 归一化没有 AutoAugment、mixup/cutmix、Random Erasing——因为主干被冻结增强只作用于特征提取太强的增强无意义甚至有害。BN 头第 222 行model.head torch.nn.Sequential(torch.nn.BatchNorm1d(model.head.in_features, affineFalse, eps1e-6), model.head)在线性头前加一个无仿射参数的 BatchNormMoCo v3 的惯例。冻结主干第 223-227 行除 head 外所有参数requires_gradFalse只训练分类头。此时可训练参数量仅约线性头一层打印number of params可见只有很小的量级。LARS 优化器第 252 行使用 util/lars.py 中的 LARSLayer-wise Adaptive Rate Scaling优化器lr 6.4 这样的超大学习率只有 LARS 能稳定收敛损失直接用torch.nn.CrossEntropyLoss()。4.4 结果对比论文TF/TPUvs 本仓库PT/GPU模型paper (TF/TPU)this repo (PT/GPU)ViT-Base68.067.8ViT-Large75.876.0ViT-Huge76.677.2本 PyTorch/GPU 代码在 ViT-Large/Huge 上取得了优于论文的结果ViT-Base 略低 0.2这很可能由 TF 与 PT 两个平台之间的系统差异如数值行为、训练细节所致。五、实践路线小结结合 README.md 的分类结果汇总ImageNet-1K 无外部数据ViT-B 83.6 / ViT-L 85.9 / ViT-H 86.9 / ViT-H448 87.8MAE 预训练权重在 ImageNet 及下游任务ImageNet-C/A/R/Sketch、iNaturalists、Places上的迁移能力已被充分验证。标准工作流为数据准备按 ImageFolder 结构组织${IMAGENET_DIR}/train与val评估验证用官方 fine-tuned 权重跑--eval复现精度确认环境正确端到端微调单机用torch.distributed.launch--accum_iter集群用submitit_finetune.py沿用官方推荐超参--global_pool、layer decay、drop_path、mixup/cutmix线性探测用submitit_linprobe.py或 main_linprobe.py 冻结主干训练线性头快速评估表征质量监控与续训TensorBoard 日志记录loss/lr以epoch_1000x为 x 轴跨 batch size 可对齐每 epoch 自动保存 checkpoint--resume支持断点续训。遇到 ViT-Huge 微调 NaN 时优先确认是否误用--cls_token改用--global_pool若仍有问题可关闭 AMP 但需接受更慢的训练速度。赞分享人工智能计算机视觉深度学习机器学习预训练微调【免费下载链接】maePyTorch implementation of MAE https//arxiv.org/abs/2111.06377项目地址https://gitcode.com/gh_mirrors/ma/mae点击查看免费下载相关推荐Candle 实现 MobileNetV4 图像分类推理从 timm 预训练权重到 Top-5 预测实战Candle 实现 MobileNetV4 图像分类推理从 timm 预训练权重到 Top 5 预测实战 本文围绕 Candle 开源仓库中 candle e人工智能大模型机器学习深度学习本地部署模型推理服务LocalAI模型微调本地训练与Fine-tuning实战指南LocalAI模型微调本地训练与Fine tuning实战指南 ? 痛点直击为什么需要本地模型微调 你是否遇到过这样的困境 数据隐私担忧 敏感业务数据人工智能大模型模型推理服务本地部署LLM 网关多模态AI AgentRAGMCP 服务使用 PyTorch 训练与预测 ConvNeXt 图像分类模型从数据集准备到权重微调的完整实践指南使用 PyTorch 训练与预测 ConvNeXt 图像分类模型从数据集准备到权重微调的完整实践指南 本篇技术指南以 pytorch_classificati示例工程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表