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")

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account