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

Go deeper with CFSGPT

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

Create Free Account