Multi-Task Learning Shared Representations Task-Specific Heads
# Multi-Task Learning: Shared Representations & Task-Specific Heads
## Introduction & Motivation
Multi-task learning shares representations across related tasks, improving generalization and sample efficiency. Soft parameter sharing: shared hidden layers + task-specific heads. Hard parameter sharing: task-specific layers share subset of weights. Foundation for transfer learning, multi-modal learning, low-resource languages.
Motivation: Tasks with limited data benefit from related tasks' signal. Shared representation avoids redundant learning. Inductive bias from multiple tasks.
Applications: Natural language processing (POS + NER + parsing), computer vision (detection + segmentation + classification), recommendation (CTR + rating prediction).
---
## Core Concepts & Theory
### Soft Parameter Sharing
Shared encoder f_shared(x) → multiple task-specific decoders g_i(f_shared(x)).
Regularize task-specific params to stay close: ||θ_i - θ_j||² penalty.
### Hard Parameter Sharing
Lower layers shared; upper layers task-specific.
Implicit regularization: fewer task-specific params.
### Task Weighting
Different loss scales per task; learnable weights λ_i optimize meta-loss.
---
## Mathematical Formulation
Multi-task objective:
$$\mathcal{L} = \sum_{i=1}^T \lambda_i \mathcal{L}_i(\hat{y}_i, y_i)$$
where λ_i ≥ 0, task weight.
Soft sharing with regularization:
$$\mathcal{L} = \sum_{i=1}^T \mathcal{L}_i + \alpha \sum_{i < j} \| heta_i - heta_j\|^2$$
Uncertainty weighting (learnable λ):
$$\mathcal{L} = \sum_{i=1}^T \frac{1}{2\sigma_i^2} \mathcal{L}_i + \sum_{i} \log \sigma_i$$
---
## Advanced Theory & Extensions
### Attention-Based Task Weighting
Learn task importance dynamically: α_i(x) per example.
### Domain-Specific Batch Normalization
Per-task batch statistics; avoid negative transfer.
### Meta-Learning for Task Adaptation
Learn initialization for rapid few-shot adaptation on new tasks.
---
## Computational Considerations
Shared encoder: O(shared_params × batch_size).
Task-specific heads: O(Σ task_params × batch_size).
Loss aggregation: O(T) tasks.
Gradient routing: Careful backprop; avoid task interference.
---
## Practical Implementation Strategies
### Task Balancing
Start with equal weights λ_i = 1/T; adjust if tasks diverge.
Gradient norm balancing: normalize gradients by per-task gradient norm.
### Negative Transfer Prevention
Monitor per-task validation loss; drop harmful tasks.
Domain-specific batch norm; per-task dropout.
### Task Scheduling
Curriculum learning: start with easier tasks; add harder tasks.
---
## Benchmark Datasets & Evaluation
NLP: Penn TreeBank (POS), CoNLL (NER), SRL.
Vision: CIFAR-10 (classification), COCO (detection + segmentation).
Metrics: Per-task accuracy/F1; average; negative transfer rate.
---
## Key Challenges & Limitations
### Negative Transfer
Tasks conflict; shared representation hurts. Solution: selective sharing, domain-specific layers.
### Imbalanced Task Data
Some tasks have more/better labels; dominate optimization.
### Task Relationship Unknown
Which tasks benefit from sharing? Requires empirical exploration.
---
## Hyperparameter Tuning
Task weights λ: 1/T initially; grid search or learnable.
Regularization α: 0.001-0.1; higher = more constraint.
Shared layer ratio: 50-90% shared; 10-50% task-specific.
---
## Real-World Applications & Case Studies
Google BERT: Masked language modeling + next sentence prediction → transfer to downstream.
Facebook Multitask Learning: Joint learning CTR + rating + recommendations.
Machine Translation: Multiple language pairs; shared encoder/decoder.
---
## Integration with Other Methods
MTL + Transfer Learning → pretrain multi-task; finetune single task.
MTL + Meta-Learning → fast adaptation to new tasks.
---
## Summary & Key Takeaways
Multi-task learning shares representations across related tasks, improving generalization via soft/hard parameter sharing and task-specific adaptation.
Principles:
1. Soft sharing: shared encoder + task-specific decoders.
2. Hard sharing: lower layers shared; upper layers task-specific.
3. Task weighting balances loss scales; learnable weights improve.
4. Negative transfer possible; monitor per-task performance.
5. Domain-specific layers prevent interference.
---
---
## Appendix: Practical Labs
### Lab 1: Soft Parameter Sharing
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
class MultiTaskNet(nn.Module):
def __init__(self, input_dim=10, shared_dim=64, task_dims=[2, 3]):
super().__init__()
self.shared = nn.Sequential(
nn.Linear(input_dim, shared_dim),
nn.ReLU()
)
# Task-specific heads
self.task_heads = nn.ModuleList([
nn.Linear(shared_dim, out_dim) for out_dim in task_dims
])
def forward(self, x):
shared_rep = self.shared(x)
outputs = [head(shared_rep) for head in self.task_heads]
return outputs
# Data
X = torch.randn(100, 10)
y1 = torch.randint(0, 2, (100,))
y2 = torch.randint(0, 3, (100,))
model = MultiTaskNet(input_dim=10, shared_dim=64, task_dims=[2, 3])
optimizer = optim.Adam(model.parameters(), lr=0.001)
criterion = nn.CrossEntropyLoss()
losses = []
for epoch in range(20):
outputs = model(X)
loss = criterion(outputs[0], y1) + criterion(outputs[1], y2)
optimizer.zero_grad()
loss.backward()
optimizer.step()
losses.append(loss.item())
print(f"Final loss: {losses[-1]:.4f}")
assert len(losses) == 20, "Should have 20 loss values"
assert all(np.isfinite(l) for l in losses), "All losses should be finite"
print("✓ Soft parameter sharing working")
if __name__ == "__main__":
print("Lab 1: Soft Sharing - PASSED")### Lab 2: Hard Parameter Sharing
import torch
import torch.nn as nn
import numpy as np
class HardSharedMTL(nn.Module):
def __init__(self, input_dim=10, hidden_dim=32, task_outputs=[2, 3]):
super().__init__()
# Shared layers
self.shared = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU()
)
# Task-specific layers
self.task_layers = nn.ModuleList([
nn.Linear(hidden_dim, out_dim) for out_dim in task_outputs
])
def forward(self, x):
shared = self.shared(x)
return [layer(shared) for layer in self.task_layers]
X = torch.randn(50, 10)
model = HardSharedMTL(input_dim=10, hidden_dim=32, task_outputs=[2, 3])
outputs = model(X)
print(f"Task 1 output shape: {outputs[0].shape}")
print(f"Task 2 output shape: {outputs[1].shape}")
assert outputs[0].shape == (50, 2), "Task 1 should output 2 classes"
assert outputs[1].shape == (50, 3), "Task 2 should output 3 classes"
print("✓ Hard parameter sharing working")
if __name__ == "__main__":
print("Lab 2: Hard Sharing - PASSED")### Lab 3: Learnable Task Weighting
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
class MTLWithWeighting(nn.Module):
def __init__(self, n_tasks=2):
super().__init__()
self.log_sigma = nn.Parameter(torch.zeros(n_tasks))
def forward(self, losses):
# Uncertainty weighting
weighted = 0
for i, loss in enumerate(losses):
weighted += 1/(2*torch.exp(self.log_sigma[i])**2) * loss + self.log_sigma[i]
return weighted
# Simulate tasks
losses_task1 = torch.tensor(2.0, requires_grad=True)
losses_task2 = torch.tensor(0.5, requires_grad=True)
mtl = MTLWithWeighting(n_tasks=2)
optimizer = optim.Adam(mtl.parameters(), lr=0.01)
for epoch in range(50):
total_loss = mtl([losses_task1, losses_task2])
optimizer.zero_grad()
total_loss.backward()
optimizer.step()
final_weights = torch.exp(-mtl.log_sigma)
print(f"Final task weights: {final_weights.detach().numpy()}")
assert len(final_weights) == 2, "Should have 2 weights"
assert (final_weights > 0).all(), "Weights should be positive"
print("✓ Learnable task weighting working")
if __name__ == "__main__":
print("Lab 3: Weighting - PASSED")### Lab 4: Multi-Task Evaluation
import numpy as np
from sklearn.tree import DecisionTreeClassifier
from sklearn.datasets import make_classification
def multi_task_evaluate(models, X_test, y_tests, task_names):
"""Evaluate multiple tasks"""
results = {}
for i, (model, y_test, name) in enumerate(zip(models, y_tests, task_names)):
acc = model.score(X_test, y_test)
results[name] = acc
avg_acc = np.mean(list(results.values()))
return results, avg_acc
# Create tasks
X_test = np.random.randn(50, 10)
y_test1 = np.random.randint(0, 2, 50)
y_test2 = np.random.randint(0, 3, 50)
# Dummy trained models
models = [
DecisionTreeClassifier(max_depth=3, random_state=0),
DecisionTreeClassifier(max_depth=3, random_state=1)
]
for model, y in zip(models, [y_test1, y_test2]):
X_train = np.random.randn(100, 10)
y_train = y[:100] if len(y) >= 100 else np.random.randint(0, 3, 100)
model.fit(X_train, y_train[:len(X_train)])
results, avg_acc = multi_task_evaluate(models, X_test, [y_test1, y_test2], ['Task1', 'Task2'])
print(f"Results: {results}, Average: {avg_acc:.2%}")
assert len(results) == 2, "Should evaluate 2 tasks"
assert 0 <= avg_acc <= 1, "Accuracy should be in [0,1]"
print("✓ Multi-task evaluation working")
if __name__ == "__main__":
print("Lab 4: Evaluation - PASSED")