何恺明 MAE 掩码自编码器演进
MAE(Masked Autoencoder)是何恺明提出的自监督视觉预训练方法,通过掩码预测任务训练 ViT,仅使用 ImageNet-1K 就能达到超越有监督预训练的效果,成为视觉基础模型训练的标准范式。
核心思想
MAE 的设计理念源自 NLP 中 BERT 的掩码语言模型,但针对视觉数据做了关键适配:
- 高掩码比例:随机掩码 75% 的图像 Patch(BERT 仅掩码 15%)
- 非对称编码器-解码器:编码器仅处理可见 Patch,解码器重建全部
- 轻量解码器:解码器远小于编码器,预训练成本极低
python
import torch
import torch.nn as nn
class MaskedAutoencoder(nn.Module):
"""MAE 掩码自编码器"""
def __init__(self, encoder, decoder, mask_ratio=0.75):
super().__init__()
self.encoder = encoder
self.decoder = decoder
self.mask_ratio = mask_ratio
def random_masking(self, x):
"""随机掩码"""
B, N, D = x.shape
num_keep = int(N * (1 - self.mask_ratio))
# 随机打乱顺序
noise = torch.rand(B, N, device=x.device)
ids_shuffle = torch.argsort(noise, dim=1)
ids_restore = torch.argsort(ids_shuffle, dim=1)
# 保留前 num_keep 个 Patch
ids_keep = ids_shuffle[:, :num_keep]
x_masked = torch.gather(x, 1, ids_keep.unsqueeze(-1).expand(-1, -1, D))
# 生成掩码
mask = torch.ones(B, N, device=x.device)
mask[:, :num_keep] = 0
mask = torch.gather(mask, 1, ids_restore)
return x_masked, mask, ids_restore
def forward(self, images):
# Patch 嵌入
patches = self.encoder.patch_embed(images)
# 掩码
visible_patches, mask, ids_restore = self.random_masking(patches)
# 编码器仅处理可见 Patch
latent = self.encoder(visible_patches)
# 解码器重建全部 Patch
reconstructed = self.decoder(latent, ids_restore)
# MSE 损失(仅在掩码位置计算)
loss = self.compute_loss(patches, reconstructed, mask)
return loss为什么 75% 掩码比例?
MAE 使用极高的 75% 掩码比例,远超 BERT 的 15%,这是基于视觉数据的特性:
- 图像冗余度高:相邻像素高度相关,低掩码比例的任务太简单
- 需要全局理解:高掩码比例迫使模型理解整体语义才能重建
- 计算效率:编码器仅处理 25% 的 Patch,计算量大幅降低
掩码比例实验
MAE 论文中的消融实验表明:30% 掩码比例几乎无法学到有意义的特征,60% 开始有效,75% 达到最佳平衡,超过 80% 则因信息不足导致重建质量下降。
MAE 的演进
MAE 提出后,多个后续工作进一步改进:
| 方向 | 方法 | 改进 |
|---|---|---|
| 视频MAE | VideoMAE | 时序掩码 + 90% 掩码比例 |
| 多模态MAE | BEiT-3 | 掩码预测统一图文 |
| 对比+MAE | iBOT | 联合对比学习和掩码预测 |
| 扩散+MAE | DiffMAE | 用扩散模型替代像素重建 |
预训练与微调
MAE 的标准流程:
- 预训练阶段:掩码重建任务训练,不使用标签
- 微调阶段:移除解码器,仅用编码器,有监督微调
- 线性探测:冻结编码器,仅训练分类头(评估特征质量)