| name | meta-learning |
| description | Guide complet du méta-apprentissage (learning to learn) — MAML, Reptile, proto-nets, few-shot, zero-shot, metric-based, optimization-based, model-based. En français. |
Méta-Apprentissage (Meta-Learning) — Guide Complet
Apprendre à apprendre : adaptation rapide avec peu de données, few-shot learning.
1. Le Problème Few-Shot
Taxonomie
Méta-Apprentissage
/ | \
Basé sur Basé sur Basé sur
l'optimis. les métriques les modèles
(MAML, (ProtoNets, (MANN,
Reptile) RelationNets) CNP)
| | |
Gradient Distance Mémoire
interne dans un externe
(inner loop) espace (external
appris memory)
2. MAML — Model-Agnostic Meta-Learning (Finn et al., 2017)
Principe
Implémentation
class MAML(nn.Module):
"""Model-Agnostic Meta-Learning.
Fonctionne avec n'importe quel modèle differentiable.
"""
def __init__(self, model, inner_lr=0.01, meta_lr=0.001,
inner_steps=5, first_order=False):
super().__init__()
self.model = model
self.inner_lr = inner_lr
self.meta_lr = meta_lr
self.inner_steps = inner_steps
self.first_order = first_order
self.meta_optimizer = torch.optim.Adam(model.parameters(), lr=meta_lr)
def forward(self, support_x, support_y, query_x):
"""Adaptation rapide (inner loop) sur support set,
prédiction sur query set.
Args:
support_x: (n_way * k_shot, ...)
support_y: (n_way * k_shot,)
query_x: (n_way * n_query, ...)
"""
fast_weights = {name: param.clone()
for name, param in self.model.named_parameters()}
for _ in range(self.inner_steps):
logits = self.model.functional_forward(support_x, fast_weights)
loss = F.cross_entropy(logits, support_y)
grads = torch.autograd.grad(loss, fast_weights.values(),
create_graph=not self.first_order)
fast_weights = {
name: w - .inner_lr * g
(name, w), g (fast_weights.items(), grads)
}
logits = .model.functional_forward(query_x, fast_weights)
logits
():
meta_loss =
task task_batch:
logits = .forward(task[], task[],
task[])
meta_loss += F.cross_entropy(logits, task[])
meta_loss /= (task_batch)
.meta_optimizer.zero_grad()
meta_loss.backward()
.meta_optimizer.step()
meta_loss.item()
(nn.Module):
():
().__init__()
.fc1 = nn.Linear(input_dim, hidden)
.fc2 = nn.Linear(hidden, hidden)
.fc3 = nn.Linear(hidden, n_way)
():
x = F.relu(.fc1(x))
x = F.relu(.fc2(x))
.fc3(x)
():
x = F.linear(x, weights[], weights[])
x = F.relu(x)
x = F.linear(x, weights[], weights[])
x = F.relu(x)
x = F.linear(x, weights[], weights[])
x
FOMAML — First-Order MAML
3. Reptile (Nichol et al., 2018)
class Reptile:
"""Reptile — meta-learning par interpolation simple.
θ ← θ + ε · (θ'_tache - θ)
Avantages :
- Pas de dérivée seconde
- Pas de forward fonctionnel
- Aussi performant que MAML
- Plus rapide
"""
def __init__(self, model, inner_lr=0.01, meta_lr=0.1, inner_steps=5):
self.model = model
self.inner_lr = inner_lr
self.meta_lr = meta_lr
self.inner_steps = inner_steps
self.optimizer = torch.optim.SGD(model.parameters(), lr=meta_lr)
def meta_train_step(self, task_batch):
for task in task_batch:
old_weights = [p.clone() for p in self.model.parameters()]
inner_optim = torch.optim.SGD(self.model.parameters(),
lr=self.inner_lr)
for _ in range(self.inner_steps):
logits = self.model(task['support_x'])
loss = F.cross_entropy(logits, task['support_y'])
inner_optim.zero_grad()
loss.backward()
inner_optim.step()
p, old_p (.model.parameters(), old_weights):
p.data = old_p + .meta_lr * (p.data - old_p)
4. ProtoNet — Prototypical Networks (Snell et al., 2017)
class ProtoNet(nn.Module):
"""Prototypical Networks for Few-Shot Learning.
Pour N-way K-shot :
1. Encoder les K exemples de chaque classe → prototypes
2. Query → encoder → distance aux prototypes
3. Softmax sur les distances négatives
Formule : P(y=c | x) = softmax(-d(f(x), p_c))
où p_c = 1/K · Σ_{x_i ∈ support_c} f(x_i)
"""
def __init__(self, encoder, distance='euclidean'):
super().__init__()
self.encoder = encoder
self.distance = distance
def forward(self, support_x, support_y, query_x):
"""support_x: (n_way * k_shot, C, H, W)
support_y: (n_way * k_shot,)
query_x: (n_way * n_query, C, H, W)
"""
support_emb = self.encoder(support_x)
query_emb = self.encoder(query_x)
n_way = len(torch.unique(support_y))
prototypes = []
for c in range(n_way):
mask = (support_y == c)
proto = support_emb[mask].mean(dim=0)
prototypes.append(proto)
prototypes = torch.stack(prototypes)
if self.distance == :
dists = torch.cdist(query_emb, prototypes)
.distance == :
query_norm = F.normalize(query_emb, dim=)
proto_norm = F.normalize(prototypes, dim=)
dists = -query_norm @ proto_norm.t()
logits = -dists
logits
():
F.cross_entropy(logits, query_y)
():
model = ProtoNet(encoder)
optimizer = torch.optim.Adam(model.parameters(), lr=)
epoch (epochs):
batch train_loader:
logits = model(batch[], batch[],
batch[])
loss = proto_loss(logits, batch[])
optimizer.zero_grad()
loss.backward()
optimizer.step()
5. Relation Networks (Sung et al., 2018)
class RelationNet(nn.Module):
"""Relation Network : apprend la métrique de similarité."""
def __init__(self, encoder, relation_module):
super().__init__()
self.encoder = encoder
self.relation_module = relation_module
def forward(self, support_x, support_y, query_x):
support_emb = self.encoder(support_x)
n_way = len(torch.unique(support_y))
prototypes = torch.stack([
support_emb[support_y == c].mean(dim=0)
for c in range(n_way)
])
query_emb = self.encoder(query_x)
n_query = query_emb.size(0)
query_expanded = query_emb.unsqueeze(1).expand(-1, n_way, -1)
proto_expanded = prototypes.unsqueeze(0).expand(n_query, -1, -1)
pairs = torch.cat([query_expanded, proto_expanded], dim=-1)
pairs = pairs.view(n_query * n_way, -1)
relations = self.relation_module(pairs).view(n_query, n_way)
return relations
6. Métriques et Evaluation
class FewShotEvaluator:
"""Évaluation few-shot avec intervalle de confiance."""
def __init__(self, model, n_tasks=1000, n_way=5, k_shot=1, n_query=15):
self.model = model
self.n_tasks = n_tasks
self.n_way = n_way
self.k_shot = k_shot
self.n_query = n_query
@torch.no_grad()
def evaluate(self, dataset):
"""Évalue sur n_tasks aléatoires."""
accuracies = []
for _ in range(self.n_tasks):
task = dataset.sample_task(self.n_way, self.k_shot, self.n_query)
logits = self.model(task['support_x'], task['support_y'],
task['query_x'])
preds = logits.argmax(dim=-1)
acc = (preds == task['query_y']).float().mean().item()
accuracies.append(acc)
mean_acc = np.mean(accuracies)
ci = 1.96 * np.std(accuracies) / np.sqrt(len(accuracies))
return {
'accuracy': mean_acc,
'ci_95': ci,
'tasks': .n_tasks,
}
7. Meta-Learning pour la Robotique
8. Meta-Learning pour les LLM (2024-2025)
class MetaPromptTuning(nn.Module):
"""Soft prompt tuning = meta-learning."""
def __init__(self, llm, n_prompt_tokens=16):
super().__init__()
self.llm = llm
self.soft_prompts = nn.Parameter(
torch.randn(1, n_prompt_tokens, llm.config.d_model)
)
def meta_learn_prompts(self, task_batch):
"""Méta-apprendre des soft prompts."""
pass
9. Algorithmes Comparés
| Algorithme | Type | 5-way 1-shot | 5-way 5-shot | Vitesse |
|---|
| MAML | Optimisation | 48.7% | 63.1% | ★★☆☆☆ |
| FOMAML | Optimisation | 48.1% | 62.8% | ★★★☆☆ |
| Reptile | Optimisation | 49.7% | 63.8% | ★★★★☆ |
| ProtoNet | Métrique | 49.4% | 68.2% | ★★★★★ |
| RelationNet | Métrique | 50.4% | 66.3% | ★★★★☆ |
| MatchingNet | Métrique | 43.6% | 56.0% | ★★★★★ |
| MetaOptNet | Optimisation | 52.2% | 69.6% | ★★★☆☆ |
Benchmark MiniImageNet
10. Implémentation Complète Few-Shot
class FewShotTask:
"""Échantillonner une tâche few-shot depuis un dataset."""
def __init__(self, dataset, n_way=5, k_shot=1, n_query=15):
self.dataset = dataset
self.n_way = n_way
self.k_shot = k_shot
self.n_query = n_query
def sample(self):
images, labels = self.dataset
classes = torch.unique(labels)
selected_classes = classes[torch.randperm(len(classes))[:self.n_way]]
support_x, support_y = [], []
query_x, query_y = [], []
for i, cls in enumerate(selected_classes):
cls_indices = (labels == cls).nonzero().squeeze(-1)
perm = torch.randperm(len(cls_indices))
support_idx = cls_indices[perm[:self.k_shot]]
support_x.append(images[support_idx])
support_y.append(torch.full((self.k_shot,), i))
query_idx = cls_indices[perm[self.k_shot:self.k_shot + self.n_query]]
query_x.append(images[query_idx])
query_y.append(torch.full((self.n_query,), i))
return {
'support_x': torch.cat(support_x),
'support_y': torch.cat(support_y),
: torch.cat(query_x),
: torch.cat(query_y),
}
Références