自监督学习:MAE、DINO 与视觉表征学习
1. 引言
自监督学习(Self-Supervised Learning, SSL)无需人工标注,直接从数据本身构造监督信号。在 ImageNet 上,自监督预训练的模型已接近甚至超越有监督预训练。
核心范式:
有监督:图像 + 标签 → 模型
自监督:图像 → 预训练任务(代理任务)→ 通用表征 → 下游任务微调
主流方法:
| 对比学习 | SimCLR, MoCo | 相似样本靠近,不同样本远离 |
| 掩码建模 | MAE, BEiT | 遮住部分输入,预测被遮部分 |
| 自蒸馏 | DINO, DINOv2 | 学生网络模仿教师网络输出 |
2. SimCLR
2.1 框架
同一图像 → 两种随机增强 → 两个视图 → 编码器 → 投影头 → 对比损失
正样本对:同一图像的两个增强视图
负样本对:不同图像的视图
损失:NT-Xent(归一化温度缩放交叉熵)
2.2 实现
import torch
import torch.nn as nn
import torchvision.transforms as T
class SimCLR(nn.Module):
def __init__(self, backbone, projection_dim=128, temperature=0.5):
super().__init__()
self.backbone = backbone
self.projector = nn.Sequential(
nn.Linear(backbone.output_dim, 512),
nn.ReLU(),
nn.Linear(512, projection_dim),
)
self.temperature = temperature
def forward(self, x1, x2):
"""x1, x2: 同一图像的两个增强视图"""
h1 = self.backbone(x1)
h2 = self.backbone(x2)
z1 = self.projector(h1)
z2 = self.projector(h2)
return z1, z2
def nt_xent_loss(self, z1, z2):
"""NT-Xent 对比损失"""
B = z1.size(0)
z = torch.cat([z1, z2], dim=0) # (2B, D)
z = nn.functional.normalize(z, dim=1)
# 相似度矩阵
sim = z @ z.T / self.temperature # (2B, 2B)
# 正样本对索引
labels = torch.cat([
torch.arange(B, 2*B),
torch.arange(0, B),
]).to(z.device)
# 排除自身
mask = torch.eye(2*B, dtype=torch.bool, device=z.device)
sim.masked_fill_(mask, –1e9)
loss = nn.CrossEntropyLoss()(sim, labels)
return loss
# 数据增强
transform = T.Compose([
T.RandomResizedCrop(224),
T.RandomHorizontalFlip(),
T.ColorJitter(0.4, 0.4, 0.4, 0.1),
T.RandomGrayscale(p=0.2),
T.GaussianBlur(kernel_size=23, sigma=(0.1, 2.0)),
T.ToTensor(),
T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
])
3. MoCo(Momentum Contrast)
3.1 核心改进
SimCLR 问题:需要大 batch(4096+)才能有足够负样本
MoCo 解决方案:
1. 动量编码器(Momentum Encoder):缓慢更新的教师网络
2. 队列(Queue):维护大量负样本(65536个)
3. 不需要大 batch
3.2 实现
class MoCo(nn.Module):
def __init__(self, backbone, dim=128, queue_size=65536, momentum=0.999, temperature=0.07):
super().__init__()
self.encoder_q = backbone # 查询编码器
self.encoder_k = backbone # 键编码器(动量更新)
self.projector_q = nn.Linear(backbone.output_dim, dim)
self.projector_k = nn.Linear(backbone.output_dim, dim)
# 初始化键编码器
for param_k in self.encoder_k.parameters():
param_k.requires_grad = False
for param_k in self.projector_k.parameters():
param_k.requires_grad = False
# 负样本队列
self.register_buffer("queue", torch.randn(dim, queue_size))
self.queue = nn.functional.normalize(self.queue, dim=0)
self.register_buffer("queue_ptr", torch.zeros(1, dtype=torch.long))
self.momentum = momentum
self.temperature = temperature
@torch.no_grad()
def _momentum_update(self):
"""动量更新键编码器"""
for param_q, param_k in zip(self.encoder_q.parameters(),
self.encoder_k.parameters()):
param_k.data = param_k.data * self.momentum + param_q.data * (1. – self.momentum)
for param_q, param_k in zip(self.projector_q.parameters(),
self.projector_k.parameters()):
param_k.data = param_k.data * self.momentum + param_q.data * (1. – self.momentum)
@torch.no_grad()
def _dequeue_and_enqueue(self, keys):
"""更新队列"""
batch_size = keys.size(0)
ptr = int(self.queue_ptr)
self.queue[:, ptr:ptr+batch_size] = keys.T
self.queue_ptr[0] = (ptr + batch_size) % self.queue.size(1)
def forward(self, x_q, x_k):
# 查询
q = self.projector_q(self.encoder_q(x_q))
q = nn.functional.normalize(q, dim=1)
# 键(动量更新)
with torch.no_grad():
self._momentum_update()
k = self.projector_k(self.encoder_k(x_k))
k = nn.functional.normalize(k, dim=1)
# 正样本相似度
l_pos = (q * k).sum(dim=1, keepdim=True) # (B, 1)
# 负样本相似度
l_neg = q @ self.queue # (B, K)
# logits
logits = torch.cat([l_pos, l_neg], dim=1) / self.temperature
labels = torch.zeros(logits.size(0), dtype=torch.long, device=logits.device)
# 更新队列
self._dequeue_and_enqueue(k)
return nn.CrossEntropyLoss()(logits, labels)
4. MAE(Masked Autoencoder)
4.1 核心思想
MAE:遮住 75% 的图像 patch,让模型重建被遮部分
编码器:只处理可见的 25% patch(高效!)
解码器:接收所有 patch(可见 + 掩码),重建原始图像
4.2 实现
class MAE(nn.Module):
def __init__(self, encoder, decoder_dim=512, patch_size=16,
mask_ratio=0.75, image_size=224):
super().__init__()
self.encoder = encoder
self.patch_size = patch_size
self.mask_ratio = mask_ratio
num_patches = (image_size // patch_size) ** 2
encoder_dim = encoder.embed_dim
# 解码器
self.decoder_embed = nn.Linear(encoder_dim, decoder_dim)
self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
self.decoder_pos_embed = nn.Parameter(
torch.zeros(1, num_patches + 1, decoder_dim)
)
self.decoder_blocks = nn.ModuleList([
TransformerBlock(decoder_dim, num_heads=8)
for _ in range(4)
])
self.decoder_norm = nn.LayerNorm(decoder_dim)
self.decoder_pred = nn.Linear(
decoder_dim, patch_size ** 2 * 3
)
def patchify(self, imgs):
"""图像 → patch"""
p = self.patch_size
B, C, H, W = imgs.shape
h, w = H // p, W // p
x = imgs.reshape(B, C, h, p, w, p)
x = x.permute(0, 2, 4, 3, 5, 1).reshape(B, h * w, p * p * C)
return x
def random_masking(self, x, mask_ratio):
"""随机遮掩"""
B, N, D = x.shape
len_keep = int(N * (1 – 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)
# 保留前 len_keep 个
ids_keep = ids_shuffle[:, :len_keep]
x_visible = torch.gather(
x, dim=1, index=ids_keep.unsqueeze(–1).expand(–1, –1, D)
)
# 生成掩码:0=保留, 1=遮掩
mask = torch.ones(B, N, device=x.device)
mask[:, :len_keep] = 0
mask = torch.gather(mask, dim=1, index=ids_restore)
return x_visible, mask, ids_restore
def forward(self, imgs):
# Patch 嵌入
x = self.encoder.patch_embed(imgs)
x = x + self.encoder.pos_embed[:, 1:, :]
# 随机遮掩
x_visible, mask, ids_restore = self.random_masking(x, self.mask_ratio)
# 编码器(只处理可见 patch)
for blk in self.encoder.blocks:
x_visible = blk(x_visible)
x_visible = self.encoder.norm(x_visible)
# 解码器
x = self.decoder_embed(x_visible)
# 填入掩码 token
mask_tokens = self.mask_token.repeat(x.size(0), ids_restore.size(1) – x.size(1), 1)
x_full = torch.cat([x, mask_tokens], dim=1)
x_full = torch.gather(
x_full, dim=1,
index=ids_restore.unsqueeze(–1).expand(–1, –1, x.size(2))
)
x = x_full + self.decoder_pos_embed[:, 1:, :]
for blk in self.decoder_blocks:
x = blk(x)
x = self.decoder_norm(x)
# 预测像素
pred = self.decoder_pred(x)
# 计算损失(只对掩码区域)
target = self.patchify(imgs)
loss = (pred – target) ** 2
loss = loss.mean(dim=–1) # per-patch loss
loss = (loss * mask).sum() / mask.sum()
return loss, pred, mask
5. DINO(Self-Distillation with No Labels)
5.1 核心思想
DINO:学生网络学习模仿教师网络的输出
教师网络 = 学生网络的指数移动平均(EMA)
输入:全局视图(224)+ 局部视图(96)
目标:局部视图的学生输出 → 匹配全局视图的教师输出
5.2 实现
class DINO(nn.Module):
def __init__(self, backbone, out_dim=65536, momentum=0.996):
super().__init__()
self.student = backbone
self.teacher = copy.deepcopy(backbone)
self.student_head = nn.Sequential(
nn.Linear(backbone.output_dim, 2048),
nn.GELU(),
nn.Linear(2048, out_dim),
)
self.teacher_head = nn.Sequential(
nn.Linear(backbone.output_dim, 2048),
nn.GELU(),
nn.Linear(2048, out_dim),
)
# 冻结教师
for p in self.teacher.parameters():
p.requires_grad = False
for p in self.teacher_head.parameters():
p.requires_grad = False
self.momentum = momentum
self.center = nn.Parameter(torch.zeros(out_dim))
@torch.no_grad()
def update_teacher(self):
"""EMA 更新教师"""
for ps, pt in zip(self.student.parameters(), self.teacher.parameters()):
pt.data = pt.data * self.momentum + ps.data * (1 – self.momentum)
for ps, pt in zip(self.student_head.parameters(), self.teacher_head.parameters()):
pt.data = pt.data * self.momentum + ps.data * (1 – self.momentum)
def forward(self, global_views, local_views):
# 教师(只处理全局视图)
with torch.no_grad():
teacher_out = [
self.teacher_head(self.teacher(v)) for v in global_views
]
teacher_out = [t – self.center for t in teacher_out]
# 学生(处理所有视图)
all_views = global_views + local_views
student_out = [
self.student_head(self.student(v)) for v in all_views
]
# 损失:学生匹配教师
loss = 0
for t_idx, t in enumerate(teacher_out):
for s_idx, s in enumerate(student_out):
if s_idx == t_idx:
continue # 跳过同一视图
loss += –(t.softmax(dim=–1) * s.log_softmax(dim=–1)).sum(dim=–1).mean()
# 更新中心
self.center = self.center * 0.9 + torch.cat(teacher_out).mean(dim=0) * 0.1
return loss / (len(teacher_out) * (len(student_out) – 1))
6. 方法对比
| SimCLR | 长 | 好 | 需要大 batch |
| MoCo | 中 | 好 | 队列存负样本 |
| MAE | 短 | 很好 | 75% 掩码,高效 |
| DINO | 中 | 最好 | 自蒸馏,无需负样本 |
| DINOv2 | 长 | SOTA | 最强视觉表征 |
7. 总结
自监督学习的核心:


