multi-task learning shared representation task transfer
# Multi-Task Learning: Shared Representation & Task Transfer
## Introduction & Motivation
Multi-Task Learning: learn multiple related tasks jointly. Share representations; improve generalization. Hard and soft parameter sharing. Applications: robotics, NLP (POS, NER), computer vision.
Motivation: Leverage task relationships; reduce overfitting.
Applications: Multiple related tasks, knowledge sharing.
---
## Core Concepts & Theory
### Hard Parameter Sharing
Shared hidden layers; task-specific heads.
### Soft Parameter Sharing
Regularized task-specific parameters; similarity.
### Cross-Task Transfer
Knowledge exchange between tasks.
---
## Mathematical Formulation
Multi-task loss:
$$L = \sum_t w_t L_t(f_{ ext{shared}}(x), y_t)$$
Task-specific loss:
$$L_t = L_{ ext{task}}(f_t(f_{ ext{shared}}(x)), y_t)$$
Regularization (soft sharing):
$$L_{ ext{reg}} = \lambda \sum_{t, t'} \|W_t - W_{t'}\|^2$$
---
## Advanced Theory & Extensions
### Attention-based Task Weighting
Dynamic task weight learning.
### Task-Attention Networks
Learn task-specific feature selection.
### Uncertainty Weighting
Automatic task balancing; learned uncertainty.
---
## Computational Considerations
Shared encoding: O(shared_model_size).
Task heads: O(sum of task sizes).
Training: Linear in number of tasks.
---
## Practical Implementation Strategies
### Task Weighting
Weight loss by task importance/difficulty.
### Early Stopping
Monitor on multiple tasks.
### Auxiliary Tasks
Use related tasks for regularization.
---
## Benchmark Datasets & Evaluation
MTL-Text: NLP multi-task benchmark.
NYUv2: Depth, surface normal, semantic segmentation.
Cityscapes: Multiple vision tasks.
---
## Key Challenges & Limitations
### Task Balancing
Different task scales; hyperparameter tuning.
### Negative Transfer
One task hurts another.
### Architecture Design
How much sharing; task-specific layers.
---
## Hyperparameter Tuning
Task weights: 0.1-10 per task.
Shared layer fraction: 50-80% shared.
Learning rate: Unified or per-task.
---
## Real-World Applications & Case Studies
NLP: Sentiment + NER + QA.
Vision: Depth + Surface Normal + Semantic Seg.
Robotics: Multiple manipulation tasks.
---
## Integration with Other Methods
MTL + Meta-learning → multi-task adaptation.
MTL + Transfer → hierarchical tasks.
---
## Summary & Key Takeaways
Multi-Task Learning via shared representations enables efficient learning of related tasks through joint optimization and cross-task transfer.
Principles:
1. Parameter sharing: shared backbone.
2. Task weighting: balance scales.
3. Task-specific layers: specialization.
4. Cross-task transfer: knowledge exchange.
5. Uncertainty weighting: automatic balancing.
---
---
## Appendix: Practical Labs
### Lab 1: Multi-Task Loss
import numpy as np
def multi_task_loss(logits_list, targets_list, task_weights=None):
"""Compute multi-task learning loss"""
num_tasks = len(logits_list)
if task_weights is None:
task_weights = np.ones(num_tasks) / num_tasks
total_loss = 0
for t in range(num_tasks):
# Task loss
exp_logits = np.exp(logits_list[t] - np.max(logits_list[t], axis=1, keepdims=True))
probs = exp_logits / exp_logits.sum(axis=1, keepdims=True)
task_loss = -np.log(
probs[np.arange(len(logits_list[t])), targets_list[t]] + 1e-8
).mean()
total_loss += task_weights[t] * task_loss
return total_loss
# Test
np.random.seed(42)
logits1 = np.random.randn(32, 10)
logits2 = np.random.randn(32, 5)
targets1 = np.random.randint(0, 10, 32)
targets2 = np.random.randint(0, 5, 32)
loss = multi_task_loss([logits1, logits2], [targets1, targets2])
assert np.isfinite(loss), "Loss finite"
print("✓ Multi-task loss working")
if __name__ == "__main__":
print("Lab 1: MultiTaskLoss - PASSED")### Lab 2: Task Weighting
import numpy as np
def compute_task_weights(losses, method='inverse'):
"""Compute dynamic task weights"""
num_tasks = len(losses)
if method == 'inverse':
# Weight by inverse loss
weights = 1.0 / (np.array(losses) + 1e-8)
elif method == 'uncertainty':
# Uncertainty weighting
weights = np.exp(-np.array(losses))
elif method == 'equal':
weights = np.ones(num_tasks)
else:
raise ValueError(f"Unknown method: {method}")
# Normalize
weights = weights / weights.sum()
return weights
# Test
losses = [0.5, 0.8, 0.3]
weights = compute_task_weights(losses, method='inverse')
assert len(weights) == 3, "Weight count"
assert np.isclose(weights.sum(), 1.0), "Normalized"
print("✓ Task weighting working")
if __name__ == "__main__":
print("Lab 2: TaskWeighting - PASSED")### Lab 3: Parameter Sharing
import numpy as np
def soft_parameter_sharing_loss(w_task1, w_task2, lambda_reg=0.1):
"""Regularize task-specific parameters toward each other"""
# L2 distance between parameters
param_diff = w_task1 - w_task2
sharing_loss = np.sum(param_diff ** 2)
return lambda_reg * sharing_loss
# Test
np.random.seed(42)
w1 = np.random.randn(100)
w2 = np.random.randn(100)
loss = soft_parameter_sharing_loss(w1, w2)
assert loss >= 0, "Loss non-negative"
print("✓ Parameter sharing working")
if __name__ == "__main__":
print("Lab 3: ParameterSharing - PASSED")### Lab 4: Multi-Task Evaluation
import numpy as np
def multi_task_evaluation(predictions_list, targets_list, metrics=['accuracy']):
"""Evaluate performance on multiple tasks"""
num_tasks = len(predictions_list)
results = {}
for t in range(num_tasks):
pred = predictions_list[t]
target = targets_list[t]
if pred.ndim > 1:
pred = np.argmax(pred, axis=1)
accuracy = (pred == target).mean()
results[f'task_{t}'] = accuracy
# Average accuracy
results['avg_accuracy'] = np.mean([results[f'task_{t}'] for t in range(num_tasks)])
return results
# Test
np.random.seed(42)
preds1 = np.random.randint(0, 10, 100)
preds2 = np.random.randint(0, 5, 100)
targets1 = np.random.randint(0, 10, 100)
targets2 = np.random.randint(0, 5, 100)
results = multi_task_evaluation([preds1, preds2], [targets1, targets2])
assert 'task_0' in results, "Task 0 results"
assert 'avg_accuracy' in results, "Average accuracy"
print("✓ Multi-task evaluation working")
if __name__ == "__main__":
print("Lab 4: MultiTaskEvaluation - PASSED")