MAE自监督预训练在CIFAR10上超越监督学习:ViT图像分类实战
发布时间:2026/9/16 5:04:53 锦皓数字建站

简介基于CIFAR10的MAE掩码自编码器实现旨在验证何凯明提出的MAE预训练策略在图像识别任务中优于直接监督学习的结论。资源面向有一定深度学习基础、关注自监督学习与ViT模型训练的开发者提供完整可运行的代码与实验记录。包体包含21个文件以Python脚本模型定义、预训练与分类器训练、PyTorch权重文件pth、配置与说明文档为主压缩包约228.74MB并附有TensorBoard日志与重建效果图便于对比MAE预训练与从头训练的效果差异。目前已有1265人学习使用。通过该资源读者可获得CIFAR10数据集上的MAE预训练管线、ViT-Tiny分类器实现、预训练权重以及训练/验证脚本快速复现自监督预训练提升下游任务精度的实验过程掌握MAE关键细节与可视化分析方法整体目录结构清晰便于直接上手实践。1. 一张 32×32 的图为什么偏要遮掉四分之三再学如果只给一次训练机会你会让模型先跑一轮遮罩重建的自我监督预训练还是直接用标签去监督学习在 CIFAR10 这个五万张 32×32 小图上何凯明的 MAEMasked Autoencoder给了一个乍看反直觉的结论自监督预训练出来的 ViT在分类任务上能打赢同一个网络从头用标签训练的版本而且优势还不小。这套开源包正好把这条实验路径完整走了一遍——mae_pretrain.py负责遮罩重建预训练train_classifier.py负责两种初始化方式的对比vit-t-mae.pth和两个分类器权重文件、TensorBoard 日志都替你备好了。对想复现MAE 在 CIFAR10 上比监督学习更有效这条结论、又不想从零搭管线的人来说这套东西可以直接当实验脚手架用。下面按预训练、分类、可视化的顺序把实现拆开讲。2. MAE 的运作机制与 ViT-T 模型选型2.1 遮罩重建为什么能在小数据集上生效MAE 的核心操作很直接把输入图像切成固定大小的 patch随机遮蔽掉其中 75% 的 patch只把可见的 25% 送进 encoder再由一个轻量 decoder 在完整位置序列上重建像素。因为遮蔽比例足够高模型无法靠邻域像素的局部连续性蒙混过关必须理解这是一只猫的耳朵这类全局语义才能补出被遮住的区域。在 CIFAR10 这种小数据集上这个特性恰好变成了一种强正则化。直接从标签学ViT 很容易把高频噪声也背下来在 5 万张图上过拟合是常态而 MAE 的预测目标是被遮蔽的像素模型被迫学习可泛化的表征后面下游任务再微调时收敛速度和精度都会改善。还有一个容易忽略的点MAE 的 loss 只算被遮蔽的 patch可见 patch 不参与反向传播所以预训练的有效计算量比常规重建方法小这是它在小显存机器上也能跑起来的原因之一。2.2 ViT-T 的模型结构这个包里的 backbone 是 ViT-TTiny 版不是标准 ViT-Base。在 32×32 输入上如果沿用 ImageNet 的 patch_size16图像会被切成 2×2 共 4 个 patch信息损失太严重所以这里 patch_size 取 4得到 8×864 个 patchembedding 维度为 192Transformer 层数为 12多头注意力的头数为 3。整体参数量在 5M 左右单张 2080Ti 就能带动。model.py里的编码器核心结构如下# model.py 关键片段已精简 import torch.nn as nn class ViT_Encoder(nn.Module): def __init__(self, img_size32, patch_size4, in_chans3, embed_dim192, depth12, num_heads3): super().__init__() self.patch_embed nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, 1 (img_size // patch_size) ** 2, embed_dim)) self.blocks nn.ModuleList([ TransformerBlock(embed_dim, num_heads) for _ in range(depth) ]) def forward(self, x): B x.shape[0] x self.patch_embed(x).flatten(2).transpose(1, 2) # [B, 64, 192] x torch.cat([self.cls_token.expand(B, -1, -1), x], dim1) x x self.pos_embed for blk in self.blocks: x blk(x) return x这段代码里patch_embed用步长等于卷积核大小的 2D 卷积一步完成切块和线性投影cls_token加在序列头部分类阶段会用到但在 MAE 预训练重建时它不参与像素预测pos_embed是绝对位置编码长度是 1 64多出来的 1 给 cls token。1 (img_size // patch_size) ** 2这个写法保证了在改输入分辨率时位置编码维度能跟着自动算出来。2.3 关键超参数速查参数MAE 预训练分类器训练mask_ratio0.75无patch_size44embed_dim192192depth1212optimizerAdamWAdamW基础学习率1.5e-41e-3weight_decay0.050.05训练轮数200100batch_size128128mask_ratio 是 MAE 里最敏感的超参数。论文在 ImageNet 上验证了 75% 附近最优CIFAR10 上因为图像本身信息密度更高我曾试过 70%重建 loss 更低但下游分类掉点说明遮蔽率太低会让模型学到边缘补齐这种投机行为。下表里的学习率也需要在预训练和微调之间区分预训练用 1.5e-4 比较稳微调分类器时 backbone 已经收敛可以放宽到 1e-3。3. MAE 预训练的实现mae_pretrain.py 逐段拆解3.1 完整流程预训练循环的核心是四步生成随机掩码、编码器只看可见 patch、decoder 预测完整序列、loss 只算被遮蔽位置。下面是从mae_pretrain.py中提取的训练循环# mae_pretrain.py 训练循环核心 for batch_idx, (images, _) in enumerate(train_loader): images images.to(device) # 1. 生成随机掩码mask_ratio0.75 mask torch.rand(B, num_patches) mask_ratio # True 表示保留 mask mask.to(device) # 2. 只把可见 patch 送进编码器 visible_x gather_visible(images, mask, patch_size) # [B, visible_num, embed_dim] features encoder(visible_x) # 3. decoder 重建完整序列含被遮蔽位置 pred decoder(features, mask) # [B, num_patches, patch_dim] # 4. 目标标准化后只在遮蔽区域算 MSE target patchify(images, patch_size) # [B, num_patches, patch_area*3] target normalize_pixels(target) # 逐 patch 做标准化 loss (pred - target) ** 2 * (~mask.unsqueeze(-1)) loss loss.sum() / (~mask).sum()这里mask是布尔张量True代表保留可见 patchgather_visible根据掩码把 patch 序列压缩拼上 cls token 后送入编码器可见 patch 数量大约只有 16 个64 × 25%。两点需要特别说明第一normalize_pixels把每个 patch 内的像素减去均值除以标准差。原始 MAE 论文发现直接预测标准化后的像素重建质量更好原因是不同 patch 的亮度基线差异被消除decoder 不用花容量去学亮度偏置。我在实际测试中看到不归一化时重建图像会整体发灰。第二loss 计算时用~mask扩展维度做选择只统计被遮蔽的 48 个 patch。如果让可见 patch 也参与 loss模型会学着把可见位置的纹理复制到邻域遮蔽比例形同虚设。3.2 优化器与调度策略预训练阶段采用 AdamW 配合余弦退火学习率初始学习率 1.5e-4weight_decay 0.05。ViT 的 patch embedding 和位置编码对 weight_decay 的敏感度不同我的经验是统一设 0.05 问题不大但要避免给 bias 和 LayerNorm 参数也加权重衰减。常见做法是为这两类参数建一个分组weight_decay 设为 0。除了优化器还有几个在 CIFAR10 上经常被忽略的设置python mae_pretrain.py --epochs 200 --batch_size 128 \ --mask_ratio 0.75 --base_lr 1.5e-4 --seed 42 --amp--amp开启混合精度训练在 200 轮预训练里能省接近一半时间。注意 CIFAR10 图像小数据增强对 MAE 不是必需品原论文在预训练阶段不做随机裁剪只做了最简单的中心归一化。如果你在预训练阶段加了强增强重建目标会变得不稳定loss 曲线会抖动。3.3 预训练阶段常见的三个坑第一个坑是随机种子不固定。MAE 的掩码是随机采样的不固定 seed 的话每次跑出来的可见 patch 集合都不一致下游微调时对比实验就没有意义。建议在mae_pretrain.py和train_classifier.py里都固定 seed42。第二个坑是 LayerNorm 的位置。ViT 里有 pre-norm 和 post-norm 两种写法MAE 论文用的是 pre-norm每个 block 里先 norm 再进 attention如果换成 post-norm深层网络训练会不稳定表现为 loss 前期不降、后期崩溃。第三个坑是 checkpoint 的保存姿势。预训练结束后vit-t-mae.pth里只保存 encoder 的 state_dictdecoder 直接丢弃。decoder 结构轻巧通常只有 4 层、embed_dim 只有 192×4但留着占空间下游也用不到。保存时建议用torch.save(encoder.state_dict(), vit-t-mae.pth)加载时再按需构造。4. 分类器训练from_scratch 与 from_pretrained 的对照4.1 实验设计的目的这个包的主要目的不是刷 CIFAR10 的精度榜而是重现MAE 预训练 ViT 优于监督训练的对比。所以train_classifier.py会在完全相同的分类器结构下跑两个配置scratch-cls从随机初始化权重开始训练pretrain-cls加载vit-t-mae.pth的 encoder 权重再微调。两个配置唯一差别是 backbone 初始化分类头都是随机初始化的线性层。4.2 模型构建与权重加载# train_classifier.py 关键片段 def build_model(init_mode, weight_pathNone): model ViT_T(num_classes10) if init_mode pretrained: ckpt torch.load(weight_path, map_locationcpu) # strictFalse忽略 cls_token 和 pos_embed 维度不匹配的键 model.load_state_dict(ckpt, strictFalse) return model # 训练时两个实验用同样配置 # python train_classifier.py --mode scratch --epochs 100 --lr 1e-3 # python train_classifier.py --mode pretrained # --init ckpt/vit-t-mae.pth --epochs 100 --lr 1e-3strictFalse是有意为之。预训练 checkpoint 里 encoder 的 cls token 是随机初始化的分类器重新建了带 10 类输出的线性头这两部分在load_state_dict时会报 key 不匹配跳过即可。预训练和微调阶段的位置编码长度一致都是 65 个 token所以可以完整加载。4.3 训练策略差异微调阶段的做法和常规分类训练不太一样# 微调阶段参数分组 for name, param in model.named_parameters(): if head in name: param.requires_grad True lr_scale 10.0 # 分类头用更大的学习率 else: param.requires_grad True lr_scale 1.0论文的微调策略不是把所有层统一用一个小学习率而是对刚初始化的分类头放一个大学习率backbone 用基础学习率。原因是分类头从头学收敛速度远慢于已经预训练过的 encoder。两个实验共用这份代码所以 from_scratch 走的也是同样的学习率分组保证公平。4.4 我在复现中观察到的结果以这份权重文件在 CIFAR10 测试集上的表现来看从零监督学习大约能到 80% 出头的准确率MAE 预训练后微调普遍比它高出 5 到 8 个百分点。这个差距在训练早期就出现pretrained 模式在第一个 epoch 就能超过 50% 准确率而 scratch 到第 5 个 epoch 才勉强到 40%。这说明预训练权重提供的初始化位置离最优解更近。实验初始化方式测试集准确率参考区间收敛轮数scratch-cls随机初始化80% ~ 82%约 60 epochspretrain-clsvit-t-mae.pth86% ~ 89%约 25 epochs不同随机种子和 gpu 型号会带来 ±1% 波动实际以自己复现为准。重点看的是相对差距不是绝对数字。5. TensorBoard 可视化与权重文件验证技巧5.1 启动 TensorBoard 观察训练动态项目里logs目录已按实验做了子目录区分mae-pretrain存放预训练阶段的标量pretrain-cls和scratch-cls存放两个分类实验的标量。启动命令tensorboard --logdir logs --port 6006浏览器打开localhost:6006重点关注两个对比预训练阶段的重建 loss 曲线是否平滑下降正常应在 200 轮内从 0.3 左右降到 0.02 附近分类阶段 pretrain-cls 和 scratch-cls 的 train/valid accuracy 两条曲线的相对位置。TensorBoard 的 scalars 面板支持直接勾选两个实验对比比翻日志文件直观得多。5.2 用权重文件做重建效果自检训练完或拿到现成权重后最快验证 MAE 学没学到的办法是手动做一次重建import torch from model import MaskedAutoencoder mae MaskedAutoencoder() ckpt torch.load(vit-t-mae.pth, map_locationcpu) mae.load_state_dict(ckpt, strictFalse) mae.eval() # 读取任意一张 CIFAR10 图片标准化后送入模型 with torch.no_grad(): pred, mask mae(images) # 返回预测和掩码重建图上被遮蔽区域的轮廓应能大致还原出物体形状边缘会有些模糊但如果出现大块色斑或整片重复纹理说明预训练没有收敛需要检查 loss 是否真的降到了 0.02 以下再考虑增加遮蔽率或轮数。5.3 核对三个 pth 文件的加载细节包根目录有两个分类器权重vit-t-classifier-from_scratch.pth是完整分类模型vit-t-classifier-from_pretrained.pth是微调后的模型。验证加载是否完整的命令for name in [vit-t-mae.pth, vit-t-classifier-from_pretrained.pth]: ckpt torch.load(name, map_locationcpu) keys list(ckpt.keys()) print(name, len(keys), keys[0], tuple(ckpt[keys[0]].shape))正常输出中from_pretrained文件应包含head.weight和head.bias两个分类头参数shape 分别为[10, 192]和[10]而vit-t-mae.pth里没有以head开头的键。用这个特征可以快速判断是不是拿错了文件——很多人把预训练权重直接当分类器加载报 shape 不匹配就是卡在这一步。本文还有配套的精品资源点击获取
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。