欢迎光临
我们一直在努力

自监督学习:MAE、DINO 与视觉表征学习

自监督学习: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. 总结

自监督学习的核心:

  • 对比学习(SimCLR/MoCo):拉近正样本,推远负样本
  • 掩码建模(MAE):遮住 75%,重建原图,简单高效
  • 自蒸馏(DINO):学生模仿教师,无需负样本和大 batch
  • DINOv2 是目前最强的通用视觉表征,微调即可用于任何下游任务
  • 赞(0)
    未经允许不得转载:171主机测评 » 自监督学习:MAE、DINO 与视觉表征学习
    分享到: 更多 (0)

    评论 抢沙发

    • 昵称 (必填)
    • 邮箱 (必填)
    • 网址