Meta-Learning Learning to Learn Maml
# Meta-Learning: Learning to Learn & MAML
## Introduction & Motivation
Meta-learning: learn to adapt quickly to new tasks. MAML (Model-Agnostic Meta-Learning): gradient-based adaptation; few gradient steps on new task. Prototypical networks: metric learning; nearest prototype. Learning rate as meta-parameter: optimize adaptation speed. Applications: few-shot learning, rapid task adaptation, continual learning.
Motivation: Real-world: limited labeled data per task. Meta-learning enables rapid adaptation; learn good initialization.
Applications: Few-shot learning, domain adaptation, continual learning.
---
## Core Concepts & Theory
### MAML (Model-Agnostic Meta-Learning)
Learn initialization enabling quick adaptation. Inner loop: task-specific steps. Outer loop: meta-update.
### Prototypical Networks
Learn metric space; classify via distance to prototype per class.
### Metric Learning
Optimize embedding; make same-class close, different-class far.
---
## Mathematical Formulation
MAML Inner Loop (per task):
$$ heta_i' = heta - \alpha
abla L_i( heta)$$
task-specific gradient step.
MAML Outer Loop (meta-update):
$$ heta \leftarrow heta - \beta
abla_ heta L_{ ext{meta}} = heta - \beta
abla_ heta \sum_i L_i( heta_i')$$
update initial weights via meta-task performance.
Prototypical Networks:
$$ ext{dist}(x, c) = \|x - c\|^2, \quad c = ext{mean of class samples}$$
---
## Advanced Theory & Extensions
### FOMAML (First-Order MAML)
Approximate second-order meta-gradients; reduce computation.
### Prototypical Networks Variants
Attention-based prototypes; learned distance metrics.
### Optimization-Based Meta-Learning
RMSprop as inner optimizer; learn optimizer parameters.
---
## Computational Considerations
MAML: O(2·forward + 2·backward) per inner step; expensive.
Prototypical: O(forward + backward) similar to supervised.
Meta-gradient: O(T·inner_steps) for T tasks.
---
## Practical Implementation Strategies
### Inner Learning Rate
Typically 0.01-0.1; smaller for stability.
### Meta Learning Rate
0.001-0.01; standard SGD/Adam ranges.
### Inner Steps
1-5 typical for few-shot; balance adaptation-computation.
---
## Benchmark Datasets & Evaluation
miniImageNet: Few-shot classification standard; 5-way 1-shot.
omniglot: Character recognition; low-shot benchmark.
CUB (birds): Fine-grained; domain adaptation benchmark.
---
## Key Challenges & Limitations
### Computational Cost
Second-order gradients expensive; FOMAML approximation common.
### Task Distribution
Meta-training tasks should match test tasks.
### Overfitting to Meta-Train
Meta-overfit on train tasks; validate on held-out tasks.
---
## Hyperparameter Tuning
Inner LR α: 0.01-0.1; dataset dependent.
Meta LR β: 0.001-0.01; standard ranges.
Inner steps: 1-5; more for hard adaptation.
---
## Real-World Applications & Case Studies
Few-Shot Image Classification: 5-way 1-shot; MAML standard baseline.
Domain Adaptation: Meta-learn to adapt; new domains quickly.
Robotics: Meta-learn dynamics; adapt to new morphologies.
---
## Integration with Other Methods
Meta-Learning + Data Augmentation → improve few-shot.
Meta-Learning + Ensemble → task-specific ensembles.
---
## Summary & Key Takeaways
Meta-learning via MAML and prototypical networks enables rapid task adaptation through learned initializations and metric spaces, achieving low-shot learning.
Principles:
1. MAML: learn initialization via meta-gradient.
2. Inner loop: adapt via task-specific gradients.
3. Outer loop: meta-update on adapted parameters.
4. Prototypical: metric-based; distance to prototypes.
5. Few-shot: minimal task data; leverage meta-prior.
---
---
## Appendix: Practical Labs
### Lab 1: MAML Forward Pass
import torch
import torch.nn as nn
class MAMLModel(nn.Module):
def __init__(self, input_dim=28*28, hidden_dim=64, output_dim=5):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, output_dim)
)
def forward(self, x):
return self.net(x)
def inner_loop(model, X_task, y_task, inner_lr=0.01, steps=1):
"""MAML inner loop: adapt to task"""
params = [p for p in model.parameters() if p.requires_grad]
for _ in range(steps):
output = model(X_task)
loss = ((output - y_task) ** 2).mean()
grads = torch.autograd.grad(loss, params, create_graph=True)
# Update parameters
with torch.no_grad():
for p, g in zip(params, grads):
p.data = p.data - inner_lr * g.data
return model
# Test
model = MAMLModel()
X_task = torch.randn(4, 28*28)
y_task = torch.randn(4, 5)
adapted_model = inner_loop(model, X_task, y_task, inner_lr=0.01, steps=1)
output = adapted_model(X_task)
assert output.shape == (4, 5), "Output shape correct"
print("✓ MAML inner loop working")
if __name__ == "__main__":
print("Lab 1: MAML - PASSED")### Lab 2: Prototypical Networks
import torch
import numpy as np
def compute_prototypes(support_features, support_labels, n_classes):
"""Compute class prototypes"""
prototypes = []
for c in range(n_classes):
mask = support_labels == c
class_features = support_features[mask]
prototype = class_features.mean(dim=0)
prototypes.append(prototype)
return torch.stack(prototypes)
def prototypical_distance(query_features, prototypes):
"""Compute distance to prototypes"""
# query_features: [n_query, feature_dim]
# prototypes: [n_classes, feature_dim]
dists = torch.cdist(query_features, prototypes, p=2)
logits = -dists # Negative distance as similarity
return logits
# Test
np.random.seed(42)
support_features = torch.randn(10, 64)
support_labels = torch.tensor([0, 0, 1, 1, 2, 2, 3, 3, 4, 4])
query_features = torch.randn(8, 64)
n_classes = 5
prototypes = compute_prototypes(support_features, support_labels, n_classes)
logits = prototypical_distance(query_features, prototypes)
assert prototypes.shape == (5, 64), "Prototype shape correct"
assert logits.shape == (8, 5), "Logit shape correct"
print("✓ Prototypical networks working")
if __name__ == "__main__":
print("Lab 2: Prototypical - PASSED")### Lab 3: Meta-Gradient Computation
import torch
import torch.nn as nn
def meta_gradient_step(model, tasks, inner_lr=0.01, meta_lr=0.001):
"""Compute MAML meta-gradient"""
optimizer = torch.optim.SGD(model.parameters(), lr=meta_lr)
meta_loss = 0
for X_task, y_task in tasks:
# Inner loop
params = [p.clone() for p in model.parameters()]
for p in model.parameters():
p.data = p.data.clone().detach().requires_grad_(True)
output = model(X_task)
task_loss = ((output - y_task) ** 2).mean()
# Compute adapted weights
grads = torch.autograd.grad(task_loss, model.parameters(), create_graph=True)
# Meta-loss on adapted parameters
meta_loss += task_loss
# Meta-update
optimizer.zero_grad()
meta_loss.backward()
optimizer.step()
return meta_loss.item()
# Test
model = nn.Sequential(nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 1))
tasks = [
(torch.randn(4, 10), torch.randn(4, 1)),
(torch.randn(4, 10), torch.randn(4, 1))
]
meta_loss = meta_gradient_step(model, tasks)
assert meta_loss > 0, "Meta-loss should be positive"
assert np.isfinite(meta_loss), "Meta-loss should be finite"
print("✓ Meta-gradient working")
if __name__ == "__main__":
print("Lab 3: MetaGradient - PASSED")### Lab 4: Few-Shot Task Simulation
import numpy as np
def generate_few_shot_task(n_way, n_shot, feature_dim=64):
"""Generate synthetic few-shot task"""
support_features = []
support_labels = []
query_features = []
query_labels = []
for c in range(n_way):
# Class center
center = np.random.randn(feature_dim)
# Support samples
support = center + 0.1 * np.random.randn(n_shot, feature_dim)
support_features.append(support)
support_labels.extend([c] * n_shot)
# Query samples
query = center + 0.1 * np.random.randn(n_shot, feature_dim)
query_features.append(query)
query_labels.extend([c] * n_shot)
support_features = np.vstack(support_features)
query_features = np.vstack(query_features)
return support_features, support_labels, query_features, query_labels
# Test
np.random.seed(42)
support, s_labels, query, q_labels = generate_few_shot_task(5, 1)
assert support.shape == (5, 64), "Support shape correct"
assert query.shape == (5, 64), "Query shape correct"
assert len(s_labels) == 5, "Support labels correct"
print("✓ Few-shot task generation working")
if __name__ == "__main__":
print("Lab 4: FewShot - PASSED")