| name | multi-task-grpo-robust |
| title | Multi-Task GRPO: Reliable LLM Reasoning Across Tasks |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2602.05547 |
| keywords | ["Multi-Task Learning","Robust Optimization","GRPO","Task Reweighting","Reinforcement Learning"] |
| description | Enable balanced multi-task GRPO training via robustness-aware optimization and improvement-aware task reweighting, dynamically adjusting task weights based on both reward and loss trajectory improvement, achieving 6-28% worst-task improvements while maintaining competitive average accuracy. |
Multi-Task GRPO: Robustness-Aware Optimization Across Tasks
Applying GRPO independently to multiple tasks leads to performance imbalance where some tasks dominate training while others stagnate. Multi-Task GRPO introduces robustness-aware optimization formulated as a minimax problem, balancing average performance against task-performance disparities. Improvement-aware task reweighting combines task-level rewards with learning progress signals, preventing weight collapse while ensuring all tasks improve.
Core Concept
The key insight is that task importance depends not only on reward signals but also on whether a task is making progress. A task with zero reward but improving loss is more important than a plateau task with high reward. By jointly optimizing for average performance and worst-case performance while tracking improvement, MT-GRPO achieves both robustness and efficiency.
Architecture Overview
- Robustness-Aware Objective: Constrained optimization minimizing performance disparities while maintaining competitive average accuracy
- Improvement-Aware Task Reweighting: Combines task reward with loss trajectory to detect stagnation and allocate more training budget
- Ratio-Preserving Sampler: Addresses GRPO's unique challenge where zero-gradient tasks (identical rewards) require special handling
- Lagrangian Relaxation: Converts constrained problem into unconstrained minimax form for practical optimization
- Convergence Guarantees: Provably converges to balanced solution
Implementation
Step 1: Define Robustness-Aware Objective
Formulate multi-task optimization as constrained problem balancing robustness and average performance.
import torch
import torch.nn.functional as F
def compute_robustness_objective(task_rewards, task_weights, lambda_robustness=1.0):
"""
Robustness-aware objective combining:
1. Average performance across tasks
2. Minimum performance across tasks (worst-case)
"""
avg_reward = (task_rewards * task_weights).sum()
min_reward = task_rewards.min()
disparity = avg_reward - min_reward
robustness_loss = avg_reward - lambda_robustness * disparity
robustness_loss, avg_reward, min_reward, disparity