| name | self-supervised-learning |
| description | Guide complet de l'apprentissage auto-supervisé (SSL) — contrastive, MAE, JEPA, SimCLR, BYOL, VICReg, BERT-style MLM, DINO, I-JEPA. En français. |
Self-Supervised Learning — Guide Complet
Apprendre des représentations sans étiquettes humaines : contrastive, prédictive, non-contrastive.
1. Pourquoi Self-Supervised Learning ?
SSL
/ | \
Contrastif Génératif Prédictif
SimCLR MAE BYOL
MoCo BERT JEPA
CLIP DALL-E DINO
SwAV GPT I-JEPA
2. Apprentissage Contrastif
SimCLR (Chen et al., 2020)
class SimCLR(nn.Module):
"""Simple Framework for Contrastive Learning of Visual Representations."""
def __init__(self, encoder, proj_dim=128):
super().__init__()
self.encoder = encoder
self.projection = nn.Sequential(
nn.Linear(encoder.output_dim, 2048),
nn.ReLU(),
nn.Linear(2048, proj_dim),
)
def forward(self, x_i, x_j):
"""x_i, x_j: deux vues augmentées du même batch."""
h_i = self.encoder(x_i)
h_j = self.encoder(x_j)
z_i = F.normalize(self.projection(h_i), dim=1)
z_j = F.normalize(self.projection(h_j), dim=1)
return z_i, z_j
def nt_xent_loss(z_i, z_j, temperature=0.5):
"""Normalized Temperature-scaled Cross Entropy Loss.
Loss = -log( exp(sim(i,j)/τ) / Σ_k exp(sim(i,k)/τ) )
Où k ∈ {vues positives et négatives du batch}
"""
B = z_i.size(0)
z = torch.cat([z_i, z_j], dim=0)
sim = torch.mm(z, z.t()) / temperature
mask = torch.eye(2 * B, device=z.device).bool()
sim = sim.masked_fill(mask, -1e9)
positive = torch.cat([
torch.arange(B, 2 * B, device=z.device),
torch.arange(0, B, device=z.device),
])
loss = F.cross_entropy(sim, positive)
return loss
transform_simclr = transforms.Compose([
transforms.RandomResizedCrop(224, scale=(0.08, 1.0)),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(0.8, 0.8, 0.8, 0.2),
transforms.RandomGrayscale(p=0.2),
transforms.GaussianBlur(kernel_size=23),
transforms.ToTensor(),
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
])
MoCo v3 (Momentum Contrast, He et al., 2021)
class MoCo(nn.Module):
"""Momentum Contrast."""
def __init__(self, encoder, queue_size=65536, momentum=0.999, dim=128):
super().__init__()
self.encoder_q = encoder
self.encoder_k = copy.deepcopy(encoder)
for p in self.encoder_k.parameters():
p.requires_grad = False
self.register_buffer("queue", F.normalize(
torch.randn(queue_size, dim), dim=1))
self.register_buffer("queue_ptr", torch.zeros(1, dtype=torch.long))
@torch.no_grad()
def _momentum_update(self):
"""Momentum update du key encoder."""
for param_q, param_k in zip(self.encoder_q.parameters(),
self.encoder_k.parameters()):
param_k.data = self.momentum * param_k.data + \
(1 - self.momentum) * param_q.data
3. BYOL — Bootstrap Your Own Latent (Grill et al., 2020)
class BYOL(nn.Module):
"""Bootstrap Your Own Latent — SSL sans négatifs."""
def __init__(self, encoder, pred_dim=256, proj_dim=256, momentum=0.996):
super().__init__()
self.online_encoder = encoder
self.online_projector = MLP(encoder.output_dim, proj_dim)
self.online_predictor = MLP(proj_dim, pred_dim)
self.target_encoder = copy.deepcopy(encoder)
self.target_projector = copy.deepcopy(self.online_projector)
for p in self.target_encoder.parameters():
p.requires_grad = False
for p in self.target_projector.parameters():
p.requires_grad = False
self.momentum = momentum
def forward(self, x1, x2):
"""x1, x2: deux vues augmentées."""
z1 = self.online_predictor(self.online_projector(self.online_encoder(x1)))
z2 = self.online_predictor(self.online_projector(.online_encoder(x2)))
torch.no_grad():
h1 = .target_projector(.target_encoder(x1))
h2 = .target_projector(.target_encoder(x2))
loss = F.mse_loss(F.normalize(z1), F.normalize(h2)) + \
F.mse_loss(F.normalize(z2), F.normalize(h1))
loss /
():
param_o, param_t (.online_encoder.parameters(),
.target_encoder.parameters()):
param_t.data = .momentum * param_t.data + \
( - .momentum) * param_o.data
4. VICReg (Bardes et al., 2022)
class VICRegLoss(nn.Module):
def __init__(self, sim_coef=25, var_coef=25, cov_coef=1):
super().__init__()
self.sim_coef = sim_coef
self.var_coef = var_coef
self.cov_coef = cov_coef
def forward(self, z1, z2):
sim_loss = F.mse_loss(z1, z2)
std_z1 = torch.sqrt(z1.var(dim=0) + 1e-4)
std_z2 = torch.sqrt(z2.var(dim=0) + 1e-4)
var_loss = torch.mean(F.relu(1 - std_z1)) + torch.mean(F.relu(1 - std_z2))
z1 = z1 - z1.mean(dim=0)
z2 = z2 - z2.mean(dim=0)
cov_z1 = (z1.T @ z1) / (z1.size(0) - 1)
cov_z2 = (z2.T @ z2) / (z2.size(0) - 1)
cov_loss = (cov_z1.pow(2).sum() - cov_z1.diag().pow(2).sum()) / z1.size(1) + \
(cov_z2.pow().() - cov_z2.diag().().()) / z2.size()
(.sim_coef * sim_loss +
.var_coef * var_loss +
.cov_coef * cov_loss)
5. MAE — Masked Autoencoder (He et al., 2022)
class MAE(nn.Module):
"""Masked Autoencoder."""
def __init__(self, vit_encoder, decoder, mask_ratio=0.75):
super().__init__()
self.encoder = vit_encoder
self.decoder = decoder
self.mask_ratio = mask_ratio
def random_masking(self, x, mask_ratio):
"""Masque aléatoire des patches.
x: (B, N, D) — tous les tokens
Retourne : tokens visibles + masque
"""
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_keep = ids_shuffle[:, :len_keep]
x_masked = torch.gather(x, dim=1, index=ids_keep.unsqueeze(-1).repeat(1, 1, D))
mask = torch.ones([B, N], device=x.device)
mask[:, :len_keep] = 0
mask = torch.gather(mask, dim=1, index=ids_shuffle)
return x_masked, mask, ids_restore
def forward(self, imgs):
patches = self.patchify(imgs)
x, mask, ids_restore = .random_masking(patches, .mask_ratio)
latent = .encoder(x)
pred = .decoder(latent, ids_restore)
target = .patchify(imgs)
loss = (pred - target) **
loss = loss.mean(dim=-)
loss = (loss * mask).() / mask.()
loss
6. JEPA — Joint Embedding Predictive Architecture (LeCun et al., 2022-2024)
class I_JEPA(nn.Module):
"""Image-based Joint Embedding Predictive Architecture.
Composants :
- Context encoder : encode la région visible
- Target encoder : encode la région cible (momentum)
- Predictor : prédit l'embedding cible depuis le contexte
"""
def __init__(self, context_encoder, target_encoder, predictor):
super().__init__()
self.context_encoder = context_encoder
self.target_encoder = target_encoder
self.predictor = predictor
for p in self.target_encoder.parameters():
p.requires_grad = False
def forward(self, x):
pass
@torch.no_grad()
def update_target():
p_c, p_t (.context_encoder.parameters(),
.target_encoder.parameters()):
p_t.data = momentum * p_t.data + ( - momentum) * p_c.data
V-JEPA (Video-JEPA, 2024)
7. DINO — Self-Distillation with No Labels (Caron et al., 2021-2023)
DINO v1 (2021)
class DINO(nn.Module):
"""Self-Distillation with No Labels."""
def __init__(self, student, teacher, center_momentum=0.9):
super().__init__()
self.student = student
self.teacher = teacher
self.register_buffer('center', torch.zeros(1, teacher.output_dim))
self.center_momentum = center_momentum
def forward(self, images):
global_views, local_views = images
with torch.no_grad():
teacher_out = self.teacher(global_views)
teacher_out = teacher_out - self.center
student_out = self.student(torch.cat([global_views, local_views]))
loss = -torch.sum(F.softmax(teacher_out / 0.04, dim=-1) *
F.log_softmax(student_out / 0.1, dim=-1), dim=-1).mean()
return loss
@torch.no_grad()
():
p_s, p_t (.student.parameters(), .teacher.parameters()):
p_t.data = * p_t.data + * p_s.data
():
.center = .center_momentum * .center + \
( - .center_momentum) * teacher_out.mean(dim=, keepdim=)
DINO v2 (2023)
8. SSL pour le NLP
BERT MLM (Masked Language Model)
class MLMLoss(nn.Module):
def forward(self, logits, labels):
return F.cross_entropy(logits.view(-1, logits.size(-1)),
labels.view(-1), ignore_index=-100)
SimCSE (Gao et al., 2021)
9. CLIP (Radford et al., 2021)
class CLIP(nn.Module):
"""Contrastive Language-Image Pre-training."""
def __init__(self, image_encoder, text_encoder):
super().__init__()
self.image_encoder = image_encoder
self.text_encoder = text_encoder
self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
def forward(self, images, texts):
image_emb = F.normalize(self.image_encoder(images), dim=-1)
text_emb = F.normalize(self.text_encoder(texts), dim=-1)
logit_scale = self.logit_scale.exp()
logits_per_image = logit_scale * image_emb @ text_emb.t()
logits_per_text = logits_per_image.t()
labels = torch.arange(len(images), device=images.device)
loss_i = F.cross_entropy(logits_per_image, labels)
loss_t = F.cross_entropy(logits_per_text, labels)
return (loss_i + loss_t) / 2
10. Tableau Comparatif
| Méthode | Négatifs | Momentum | Batch req. | Type | Année | Top-1 (ImageNet) |
|---|
| SimCLR | ✓ | Non | 4096 | Contrastif | 2020 | 76.5% |
| MoCo v3 | ✓ | ✓ | 256 | Contrastif | 2021 | 76.7% |
| BYOL | ✗ | ✓ | 2048 | Prédictif | 2020 | 77.4% |
| SwAV | ✗ | Non | 256 | Clustering | 2020 | 78.1% |
| VICReg | ✗ | Non | 2048 | Régularisation | 2022 | 78.5% |
| MAE | ✗ | ✓ | 4096 | Génératif | 2022 | 80.6% |
| I-JEPA | ✗ | ✓ | 2048 | Prédictif | 2023 | 79.3% |
| DINO v2 | ✗ | ✓ | 1024 | Distillation | 2023 | 81.1% |
Références