news 2026/10/4 6:36:08

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MAE 图像分类微调实战指南:从预训练权重评估、端到端 Fine-tuning 到 Linear Probing(PyTorch 实现)
  • 人工智能
  • 计算机视觉
  • 深度学习
  • 机器学习
  • 预训练
  • 微调

【免费下载链接】mae

PyTorch implementation of MAE https//arxiv.org/abs/2111.06377

项目地址:https://gitcode.com/gh_mirrors/ma/mae
点击查看免费下载

本文基于本仓库(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-BaseViT-LargeViT-Huge
fine-tuned checkpointmae_finetuned_vit_base.pthmae_finetuned_vit_large.pthmae_finetuned_vit_huge.pth
md51b25e951f5502541f2
reference ImageNet accuracy83.66485.95286.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.731

2.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.646

2.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.584

2.4 评估流程的源码解析

从 main_finetune.py 看,评估路径的核心逻辑如下:

  1. 参数入口:--eval开关(第 137-138 行)触发"仅评估"模式;--resume指定权重路径(第 132-133 行);--model选择模型结构(第 51-52 行,默认vit_large_patch16);--batch_size默认 64(第 44-45 行),官方示例统一用 16。
  2. 模型构建:models_vit.__dict__args.model(第 227-231 行)。--nb_classes默认 1000,适配 ImageNet-1K。
  3. 权重加载与结构适配(第 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。
  4. 精度统计: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 中与微调效果直接相关的关键参数及其默认值、含义:

参数默认值说明
--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 用 100,Large/Huge 用 50)
--blr1e-3基准学习率,实际 lr = blr × 有效 batch / 256
--layer_decay0.75层间学习率衰减系数(Base 0.65,Large/Huge 0.75)
--weight_decay0.05权重衰减
--drop_path0.1DropPath 随机深度丢弃率(Base 0.1,Large 0.2,Huge 0.3)
--reprob0.25Random Erasing 概率(此处跟随 DeiT 设置)
--mixup0mixup alpha,>0 时启用
--cutmix0cutmix alpha,>0 时启用
--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属于第 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,来自原文档)

  1. 归一化像素:官方提供的预训练权重是用归一化像素损失(--norm_pix_loss)训练 1600 epoch 得到的(论文 Table 3),微调超参数因此与使用未归一化像素的默认基线略有不同。
  2. AMP 与数值行为差异:原版 MAE 是 TensorFlow+TPU 且无显式混合精度;本复现为 PyTorch+GPU 并启用 AMP,两个平台存在数值行为差异。本仓库微调统一使用--global_pool(全局平均池化);用--cls_token效果相当,但在 GPU 上微调 ViT-Huge 时有产生 NaN 的可能(TPU 上未观察到)。关闭 AMP 可以规避该问题,但训练更慢。
  3. 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 有四处关键差异:

  1. 弱增强(第 131-141 行):训练只用 RandomResizedCrop(224) + 随机水平翻转 + 归一化,没有 AutoAugment、mixup/cutmix、Random Erasing——因为主干被冻结,增强只作用于特征提取,太强的增强无意义甚至有害。
  2. BN 头(第 222 行):model.head = torch.nn.Sequential(torch.nn.BatchNorm1d(model.head.in_features, affine=False, eps=1e-6), model.head),在线性头前加一个无仿射参数的 BatchNorm(MoCo v3 的惯例)。
  3. 冻结主干(第 223-227 行):除 head 外所有参数requires_grad=False,只训练分类头。此时可训练参数量仅约线性头一层(打印number of params可见只有很小的量级)。
  4. 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-Base68.067.8
ViT-Large75.876.0
ViT-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-H@448 87.8),MAE 预训练权重在 ImageNet 及下游任务(ImageNet-C/A/R/Sketch、iNaturalists、Places)上的迁移能力已被充分验证。标准工作流为:

  1. 数据准备:按 ImageFolder 结构组织${IMAGENET_DIR}/train与val;
  2. 评估验证:用官方 fine-tuned 权重跑--eval复现精度,确认环境正确;
  3. 端到端微调:单机用torch.distributed.launch+--accum_iter,集群用submitit_finetune.py,沿用官方推荐超参(--global_pool、layer decay、drop_path、mixup/cutmix);
  4. 线性探测:用submitit_linprobe.py或 main_linprobe.py 冻结主干训练线性头,快速评估表征质量;
  5. 监控与续训: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

项目地址:https://gitcode.com/gh_mirrors/ma/mae
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/4 6:35:57

Obsidian + WorkBuddy + Gitee:构建可对话、可追溯的个人知识库

1. 为什么我要把 Obsidian、WorkBuddy 和 Gitee 拼在一起用先说结论&#xff1a;这套组合解决的核心问题只有一个——让个人知识库从“静态笔记堆”变成“能对话、能追溯、能回滚的活系统”。我用了三年 Obsidian&#xff0c;笔记攒了四千多条&#xff0c;但真正回头翻的不到百…

作者头像 李华
网站建设 2026/10/4 6:34:17

10款降AI率工具实测,专科论文怎么选才靠谱

专科论文写作进入2026年&#xff0c;降AI率成了绕不开的一关。不少学校对毕业论文的AI生成内容检测越来越严&#xff0c;重复率过了、AI检测却亮红灯的情况比比皆是。市面上号称能降AI率的工具五花八门&#xff0c;实际效果参差不齐。这篇盘点从专科生实际写作场景出发&#xf…

作者头像 李华
网站建设 2026/10/4 6:32:49

MR25H40CDF与PIC32组合:工业存储的MRAM高速读写方案

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/4 6:28:40

SPI MRAM 免擦除存储方案:MR25H40CDF 与 TM4C129 的工业实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/10/4 6:26:31

Codex CLI 从零上手:Node.js 环境准备与模型接入避坑指南

1. 从零上手 Codex CLI&#xff1a;先搞清楚它到底解决什么问题很多人第一次听到 Codex CLI&#xff0c;脑子里冒出来的第一个问题是"这不就是个命令行版的聊天工具吗"。我一开始也这么想&#xff0c;直到真正把它接进日常开发流程之后才发现&#xff0c;它和网页端对…

作者头像 李华
网站建设 2026/10/4 6:26:26

Claude Code 2.1.287 Mods 机制解析:CLI 中间件与插件行为改写实战

1. 从 2.1.287 这个版本号说起&#xff1a;Mods 到底改了什么Claude Code 更新到 2.1.287 之后&#xff0c;最值得拿出来聊的不是某个命令的小修小补&#xff0c;而是Mods这个机制的引入。简单说&#xff0c;它让插件从"只能挂载工具、加几个斜杠命令"进化到了"…

作者头像 李华