acceptodds
Under review as a conference paper at ICLR 2027

KAGrad: KL-Aware Gradient Alignment in Multitask Learning

Abstract

In gradient-based multitask learning with shared parameters, an update that decreases one task loss can increase another. Existing methods mitigate such conflicts using the Euclidean geometry of task gradients. For probabilistic models, however, gradient geometry alone does not determine local changes in output distributions. Tasks with identical gradients can have different Fisher information matrices and therefore assign different local KL costs to the same update. We propose KL-Aware Gradient Alignment (KAGrad), which incorporates this task-specific geometry into the choice of a shared update. Each task gradient and a supplied positive-definite metric define a quadratic model of predicted improvement; with the exact Fisher, its quadratic penalty is proportional to the local KL term. KAGrad maximizes the smallest fraction of each task's independently attainable model improvement, within a Euclidean ball centered at the average-loss gradient. Equivalently, it minimizes the largest normalized squared distance, in each task's supplied metric, from the update preferred by that task's local model. The optimized worst-task fraction is nonnegative if and only if the feasible set contains an update under which every local model used in the optimization predicts non-increase. When the selected direction is applied directly as a gradient step, the constraint gives average-loss descent and an stationarity bound under smoothness, a lower-bounded average loss, and a step-size bound. Because a Fisher metric need not match task-loss curvature, we also derive a sufficient condition under which the local predictions hold for actual losses after a finite step. Against CAGrad, which uses the same average-gradient constraint with a first-order criterion, KAGrad with estimated Fisher or identity metrics attains higher Cityscapes mean intersection over union ( and versus ) at a shared constraint radius, but higher relative depth error ( and versus ). On MT50, eight-term grouped KAGrad attains higher peak mean task success than eight-term CAGrad-Fast ( versus ) over checkpoints shared across runs. We separately assess the supplied metric. In quadratic tasks, matching it to loss curvature improves worst-task actual loss decrease over an identity metric; rotating the Hessians away diminishes and reverses this advantage. In learned models, diagonal Fisher estimates show no consistent benefit over identity metrics: MT10 peak mean success is versus , and ten paired Multi-Fashion+MNIST runs detect no difference. Cityscapes shows a segmentation–depth trade-off. The benefit of a curvature-matched metric in quadratic tasks does not carry over to diagonal Fisher estimates in these learned models.

open until 14 Dec 2026

est. 32% chance this paper gets accepted at ICLR 2027.

Reject 68%Accept 32%

What do you think this paper will get?

All positions stay anonymous.

Related papers

Loading the map…

Discussion (0)

Sign in to comment.