- 人工智能
- 计算机视觉
- 深度学习
- 机器学习
- 预训练
- 微调
【免费下载链接】mae
PyTorch implementation of MAE https//arxiv.org/abs/2111.06377
本文基于本仓库(Masked Autoencoders: A PyTorch Implementation,MAE 论文《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 与评估(Acc@1/Acc@5)的实现
- models_vit.py:支持 global pooling 的 ViT 模型定义(Base/Large/Huge)
- util/lr_decay.py:BEiT 风格的层间学习率衰减(layer-wise lr decay)
- util/lr_sched.py:warmup + 半周期余弦学习率调度
- util/pos_embed.py:位置编码加载与插值(适配不同输入分辨率)
- util/datasets.py:ImageNet 数据集与训练/评估数据增强管线
- submitit_finetune.py、submitit_linprobe.py:基于 submitit 的 Slurm 多节点作业提交脚本
1.2 环境与依赖
- PyTorch + CUDA GPU(本仓库基于 PyTorch + GPU 复现,原版实现为 TensorFlow + TPU)
timm==0.3.2:main_finetune.py第 26 行有硬性版本断言assert timm.__version__ == "0.3.2",请务必锁定该版本- torchvision、tensorboard
- 多节点训练需要额外安装 submitit(
pip install submitit),单节点训练则不需要
1.3 数据集目录结构
${IMAGENET_DIR}是一个包含{train, val}两个子目录的 ImageNet 目录:
${IMAGENET_DIR}/ ├── train/ # 训练集,按类别分目录(torchvision ImageFolder 格式) └── val/ # 验证集在 util/datasets.py 中,build_dataset通过datasets.ImageFolder(root, transform=transform)加载数据,root 分别为${IMAGENET_DIR}/train与${IMAGENET_DIR}/val。
二、评估官方微调权重(Evaluation)
作为微调正确性的 sanity check,首先用官方公开的 ImageNetfine-tuned权重做纯评估。下表列出了官方三个规模微调后的权重、md5 与参考精度(数据来自 FINETUNE.md):
| 项目 | ViT-Base | ViT-Large | ViT-Huge |
|---|---|---|---|
| fine-tuned checkpoint | mae_finetuned_vit_base.pth | mae_finetuned_vit_large.pth | mae_finetuned_vit_huge.pth |
| md5 | 1b25e9 | 51f550 | 2541f2 |
| reference ImageNet accuracy | 83.664 | 85.952 | 86.928 |
权重文件托管于官方公开下载地址(dl.fbaipublicfiles.com/mae/finetune/目录),文件名与上表一致,下载后请用 md5 校验文件完整性。
2.1 单 GPU 评估 ViT-Base
python main_finetune.py --eval --resume mae_finetuned_vit_base.pth --model vit_base_patch16 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出:
* Acc@1 83.664 Acc@5 96.530 loss 0.7312.2 评估 ViT-Large
python main_finetune.py --eval --resume mae_finetuned_vit_large.pth --model vit_large_patch16 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出:
* Acc@1 85.952 Acc@5 97.570 loss 0.6462.3 评估 ViT-Huge
python main_finetune.py --eval --resume mae_finetuned_vit_huge.pth --model vit_huge_patch14 --batch_size 16 --data_path ${IMAGENET_DIR}预期输出:
* Acc@1 86.928 Acc@5 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, std=2e-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-Huge,md5 分别为8cad7c、b8b06e、9bdbb0)。官方预训练权重使用归一化像素损失(--norm_pix_loss)训练 1600 epoch(论文 Table 3),因此微调超参数与默认未归一化基线略有差异。
3.1 多节点分布式训练(4 节点 × 8 GPU)
使用 submitit 在 Slurm 集群上提交作业(需先pip install submitit):
python 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 size=
batch_size(每 GPU 32)×nodes(4)× 每节点 GPU 数(8)=1024。 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×16=256 个 patch,
drop_path提高到 0.3 以增强正则。
3.4 单节点训练(1 节点 × 8 GPU,无需 submitit)
用torch.distributed.launch代替 submitit:
OMP_NUM_THREADS=1 python -m torch.distributed.launch --nproc_per_node=8 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_iter(4)× 8(GPU)=1024。--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 中与微调效果直接相关的关键参数及其默认值、含义:
| 参数 | 默认值 | 说明 |
|---|---|---|
--model | vit_large_patch16 | 模型结构:vit_base_patch16/vit_large_patch16/vit_huge_patch14 |
--batch_size | 64 | 每 GPU batch size,有效 batch = batch_size × accum_iter × GPU 数 |
--accum_iter | 1 | 梯度累积迭代数,用于在显存受限时增大有效 batch size |
--epochs | 50 | 训练轮数(Base 用 100,Large/Huge 用 50) |
--blr | 1e-3 | 基准学习率,实际 lr = blr × 有效 batch / 256 |
--layer_decay | 0.75 | 层间学习率衰减系数(Base 0.65,Large/Huge 0.75) |
--weight_decay | 0.05 | 权重衰减 |
--drop_path | 0.1 | DropPath 随机深度丢弃率(Base 0.1,Large 0.2,Huge 0.3) |
--reprob | 0.25 | Random Erasing 概率(此处跟随 DeiT 设置) |
--mixup | 0 | mixup alpha,>0 时启用 |
--cutmix | 0 | cutmix alpha,>0 时启用 |
--smoothing | 0.1 | 标签平滑系数(启用 mixup 时由 mixup 标签变换接管) |
--aa | rand-m9-mstd0.5-inc1 | AutoAugment 策略 |
--global_pool/--cls_token | global_pool 默认开 | 分类特征来源:全局平均池化 / class token |
--finetune | 空 | 预训练权重路径 |
--data_path | /datasets01/imagenet_full_size/061417/ | 数据集根目录(含 train/val) |
--nb_classes | 1000 | 分类类别数 |
--warmup_epochs | 5 | 学习率 warmup 轮数 |
--min_lr | 1e-6 | 余弦调度下界 |
--clip_grad | None | 梯度裁剪范数(默认不裁剪) |
--dist_eval | False | 分布式评估(训练时推荐开启以加速监控) |
--eval | False | 仅评估模式 |
--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属于第 i+1 层,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_transform:AutoAugment(默认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()中进行,梯度缩放使用仓库自实现的NativeScalerWithGradNormCount(util/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 是 TensorFlow+TPU 且无显式混合精度;本复现为 PyTorch+GPU 并启用 AMP,两个平台存在数值行为差异。本仓库微调统一使用
--global_pool(全局平均池化);用--cls_token效果相当,但在 GPU 上微调 ViT-Huge 时有产生 NaN 的可能(TPU 上未观察到)。关闭 AMP 可以规避该问题,但训练更慢。 - RandErase:这里跟随 DeiT 设置
--reprob 0.25,其效果小于随机方差(即对最终精度的贡献不显著)。
四、线性探测(Linear Probing)
Linear Probing 用于评估预训练表征质量:冻结主干所有参数,只训练一个线性分类头,看特征本身的线性可分性。
4.1 4 节点 × 8 GPU 训练 ViT-Base
python 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 epoch,32 张 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, affine=False, eps=1e-6), model.head),在线性头前加一个无仿射参数的 BatchNorm(MoCo v3 的惯例)。 - 冻结主干(第 223-227 行):除 head 外所有参数
requires_grad=False,只训练分类头。此时可训练参数量仅约线性头一层(打印number of params可见只有很小的量级)。 - LARS 优化器(第 252 行):使用 util/lars.py 中的 LARS(Layer-wise Adaptive Rate Scaling)优化器,lr 6.4 这样的超大学习率只有 LARS 能稳定收敛;损失直接用
torch.nn.CrossEntropyLoss()。
4.4 结果对比:论文(TF/TPU)vs 本仓库(PT/GPU)
| 模型 | paper (TF/TPU) | this repo (PT/GPU) |
|---|---|---|
| ViT-Base | 68.0 | 67.8 |
| ViT-Large | 75.8 | 76.0 |
| ViT-Huge | 76.6 | 77.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-H@448 87.8),MAE 预训练权重在 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 但需接受更慢的训练速度。
- 人工智能
- 计算机视觉
- 深度学习
- 机器学习
- 预训练
- 微调
【免费下载链接】mae
PyTorch implementation of MAE https//arxiv.org/abs/2111.06377
相关推荐
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),仅供参考