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

Go deeper with CFSGPT

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

Create Free Account