Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,12 @@ changelog does not include internal changes that do not affect the user.

## [Unreleased]

### Added

- Added `GradNorm` and `GradNormWeighting` for adaptive task-loss balancing, with an
external optimizer for the task weights and an example using the last shared layer's
gradient norms.

## [0.17.1] - 2026-09-23

### Fixed
Expand Down
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -185,6 +185,7 @@ TorchJD provides many existing aggregators from the literature, listed in the fo
| [ExcessMTL](https://torchjd.org/stable/docs/aggregation/excess_mtl#torchjd.aggregation.ExcessMTL) | [ExcessMTLWeighting](https://torchjd.org/stable/docs/aggregation/excess_mtl#torchjd.aggregation.ExcessMTLWeighting) | [Robust Multi-Task Learning with Excess Risks](https://proceedings.mlr.press/v235/he24n.html) |
| [FairGrad](https://torchjd.org/stable/docs/aggregation/fairgrad#torchjd.aggregation.FairGrad) | [FairGradWeighting](https://torchjd.org/stable/docs/aggregation/fairgrad#torchjd.aggregation.FairGradWeighting) | [Fair Resource Allocation in Multi-Task Learning](https://arxiv.org/pdf/2402.15638) |
| [GradDrop](https://torchjd.org/stable/docs/aggregation/graddrop#torchjd.aggregation.GradDrop) | - | [Just Pick a Sign: Optimizing Deep Multitask Models with Gradient Sign Dropout](https://arxiv.org/pdf/2010.06808) |
| [GradNorm](docs/source/docs/aggregation/gradnorm.rst) | [GradNormWeighting](docs/source/docs/aggregation/gradnorm.rst) | [GradNorm: Gradient Normalization for Adaptive Loss Balancing in Deep Multitask Networks](https://proceedings.mlr.press/v80/chen18a.html) |

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We can remove that so that it doesn't appear in the readme until we actually make a release (the release skill will add this row back to the readme).

| [GradVac](https://torchjd.org/stable/docs/aggregation/gradvac#torchjd.aggregation.GradVac) | [GradVacWeighting](https://torchjd.org/stable/docs/aggregation/gradvac#torchjd.aggregation.GradVacWeighting) | [Gradient Vaccine: Investigating and Improving Multi-task Optimization in Massively Multilingual Models](https://arxiv.org/pdf/2010.05874) |
| [IMTLG](https://torchjd.org/stable/docs/aggregation/imtl_g#torchjd.aggregation.IMTLG) | [IMTLGWeighting](https://torchjd.org/stable/docs/aggregation/imtl_g#torchjd.aggregation.IMTLGWeighting) | [Towards Impartial Multi-task Learning](https://www.semanticscholar.org/paper/Towards-Impartial-Multi-task-Learning-Liu-Li/45c0828baec1dd53b81f1b2635788fdf27d0792d) |
| [Krum](https://torchjd.org/stable/docs/aggregation/krum#torchjd.aggregation.Krum) | [KrumWeighting](https://torchjd.org/stable/docs/aggregation/krum#torchjd.aggregation.KrumWeighting) | [Machine Learning with Adversaries: Byzantine Tolerant Gradient Descent](https://proceedings.neurips.cc/paper/2017/file/f4b9ec30ad9f68f89b29639786cb62ef-Paper.pdf) |
Expand Down
10 changes: 10 additions & 0 deletions docs/source/docs/aggregation/gradnorm.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
:hide-toc:

GradNorm
========

.. autoclass:: torchjd.aggregation.GradNorm
:members: __call__, set_losses, balancing_loss, renormalize, reset

.. autoclass:: torchjd.aggregation.GradNormWeighting
:members: __call__, set_losses, balancing_loss, renormalize, reset
1 change: 1 addition & 0 deletions docs/source/docs/aggregation/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ Abstract base classes
excess_mtl.rst
fairgrad.rst
graddrop.rst
gradnorm.rst
gradvac.rst
imtl_g.rst
krum.rst
Expand Down
59 changes: 59 additions & 0 deletions docs/source/examples/gradnorm.rst

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We usually have method-specific examples directly in the docstring of the class of the method. Could you move it there?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Very cool example!

Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
Adaptive loss balancing with GradNorm
======================================

GradNorm adjusts task weights using their gradient norms and their progress relative to the
initial losses. The paper computes these norms over the last shared layer. The resulting weights
still apply to every model parameter, including the task heads.

This example uses :class:`~torchjd.autogram.Engine` to compute the Gramian for that layer and
:class:`~torchjd.aggregation.GradNormWeighting` to form the auxiliary balancing loss. A separate
optimizer learns the task weights. Their sum is restored to the number of tasks after each step.
Both optimizers use gradients computed with the weights from before the step.

.. testcode::

import torch
from torch.nn import Linear, MSELoss, ReLU, Sequential
from torch.optim import Adam, SGD

from torchjd.aggregation import GradNormWeighting
from torchjd.autogram import Engine

shared = Sequential(Linear(5, 4), ReLU(), Linear(4, 3))
heads = [Linear(3, 1), Linear(3, 1)]
parameters = [*shared.parameters(), *(p for head in heads for p in head.parameters())]
model_optimizer = SGD(parameters, lr=0.01)
weighting = GradNormWeighting(n_tasks=2, alpha=1.5)
weight_optimizer = Adam(weighting.parameters(), lr=0.001)
criterion = MSELoss()
engine = Engine(shared[2], batch_dim=None)

inputs = torch.randn(4, 8, 5)
targets = torch.randn(4, 8, 2)

for features, target in zip(inputs, targets):
model_optimizer.zero_grad()
weight_optimizer.zero_grad()
representation = shared(features)
losses = torch.stack([
criterion(head(representation).squeeze(1), target[:, i])
for i, head in enumerate(heads)
])
gramian = engine.compute_gramian(losses)
weighting.set_losses(losses)
weights = weighting(gramian)
weighting.balancing_loss().backward()
losses.backward(weights)
model_optimizer.step()
weight_optimizer.step()
weighting.renormalize()

The auxiliary backward affects only the task weights. The model backward uses detached weights,
so its gradients cannot alter the balancing update. Computing the Gramian on all model parameters
instead would include the task heads in the norms, which differs from the paper's choice.
The losses are already averaged over the batch, so the engine uses ``batch_dim=None``.

Save both optimizers' states along with the model and weighting ``state_dict()`` to resume
training. The weighting stores its initial losses and learned weights; the next batch must still
call ``set_losses`` and the forward. Its ``reset()`` method starts a new loss baseline and restores
unit weights. Reset the external optimizer as well when starting a new experiment.
3 changes: 3 additions & 0 deletions docs/source/examples/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@ This section contains some usage examples for TorchJD.
- :doc:`Multi-Task Learning (MTL) <mtl>` provides an example of multi-task learning where Jacobian
descent is used to optimize the vector of per-task losses of a multi-task model, using the
dedicated backpropagation function :doc:`mtl_backward <../docs/autojac/mtl_backward>`.
- :doc:`GradNorm <gradnorm>` learns task weights from the gradient norms of the last shared layer
and applies them to the whole model.
- :doc:`Instance-Wise Multi-Task Learning (IWMTL) <iwmtl>` shows how to combine multi-task learning
with instance-wise risk minimization: one loss per task and per element of the batch, using the
:doc:`autogram.Engine <../docs/autogram/engine>`.
Expand All @@ -40,6 +42,7 @@ This section contains some usage examples for TorchJD.
iwrm.rst
partial_jd.rst
mtl.rst
gradnorm.rst
iwmtl.rst
rnn.rst
monitoring.rst
Expand Down
3 changes: 3 additions & 0 deletions src/torchjd/aggregation/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
from ._excess_mtl import ExcessMTL, ExcessMTLWeighting
from ._fairgrad import FairGrad, FairGradWeighting
from ._graddrop import GradDrop
from ._gradnorm import GradNorm, GradNormWeighting
from ._gradvac import GradVac, GradVacWeighting
from ._imtl_g import IMTLG, IMTLGWeighting
from ._krum import Krum, KrumWeighting
Expand Down Expand Up @@ -80,6 +81,8 @@
"FairGrad",
"FairGradWeighting",
"GradDrop",
"GradNorm",
"GradNormWeighting",
"GradVac",
"GradVacWeighting",
"GramianWeightedAggregator",
Expand Down
239 changes: 239 additions & 0 deletions src/torchjd/aggregation/_gradnorm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,239 @@
from math import isfinite

import torch
from torch import Tensor, nn

from torchjd._mixins import Stateful
from torchjd.linalg import PSDMatrix

from ._aggregator_bases import GramianWeightedAggregator
from ._mixins import _NonDifferentiable
from ._weighting_bases import _GramianWeighting


class GradNormWeighting(_GramianWeighting, Stateful, _NonDifferentiable):
r"""
:class:`~torchjd.Stateful`
:class:`~torchjd.aggregation.Weighting` [:class:`~torchjd.linalg.PSDMatrix`] from
`GradNorm: Gradient Normalization for Adaptive Loss Balancing in Deep Multitask Networks
<https://proceedings.mlr.press/v80/chen18a.html>`_ (ICML 2018).

The trainable parameter ``weights`` starts at one for each task. Call :meth:`set_losses`
before each forward, then minimise :meth:`balancing_loss` with an external optimizer and
call :meth:`renormalize` after its step. The forward returns detached weights for the
model's backward pass. It does not update the weights or accumulate their gradients.

For a Gramian :math:`G = JJ^\top`, the balancing loss is

.. math::
L_{\mathrm{grad}} = \sum_i \left| g_i - \operatorname{stopgrad}
\left(\bar g r_i^\alpha\right) \right|, \qquad
g_i = |w_i|\sqrt{G_{ii}}, \qquad
r_i = \frac{L_i/L_i(0)}{\operatorname{mean}_j(L_j/L_j(0))}.

Here :math:`w_i` is the learned weight, :math:`L_i` the current task loss,
:math:`L_i(0)` its value at the first call to :meth:`set_losses` and :math:`\bar g`
the mean weighted gradient norm. Only ``weights`` receives gradients from this loss.

:param n_tasks: Number of tasks. Must be positive and remains fixed for this instance.
:param alpha: Non-negative strength of training-rate balancing. With ``0``, GradNorm
targets equal gradient norms. The paper uses ``1.5`` for its NYUv2 experiments.

Move the module to the model's device and dtype before creating its optimizer. Initial
losses must be finite and strictly positive; later losses may be zero. If all current
losses are zero, the balancing loss and its weight gradients are zero.

This follows Algorithm 1's direct weight updates and normalization to a sum of
``n_tasks``. LibMTL instead learns softmax logits and uses first-epoch losses as its
baseline. Choose a weight learning rate that keeps the weights non-negative.

.. testcode::

import torch
from torch.nn import Linear
from torch.optim import Adam, SGD

from torchjd.aggregation import GradNormWeighting
from torchjd.autojac import jac

model = Linear(3, 2)
weighting = GradNormWeighting(2)
model_optimizer = SGD(model.parameters(), lr=0.01)
weight_optimizer = Adam(weighting.parameters(), lr=0.001)

for features in torch.randn(4, 8, 3):
model_optimizer.zero_grad()
weight_optimizer.zero_grad()
losses = model(features).square().mean(dim=0)
jacs = jac(losses, list(model.parameters()), retain_graph=True)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

To match the paper a bit better, we could compute only the jacobians w.r.t. the last shared layer here.

J = torch.cat([j.flatten(1) for j in jacs], dim=1)
weighting.set_losses(losses)
weights = weighting(J @ J.T)
weighting.balancing_loss().backward()
losses.backward(weights)
model_optimizer.step()
weight_optimizer.step()
weighting.renormalize()

See :doc:`the GradNorm example <../../examples/gradnorm>` for computing the norms
using only the last shared layer while weighting all model parameters.
Comment on lines +78 to +79

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If we do the changes mentioned before, this can be simply removed. We could just have 2 examples in GradNormWeighting: one using autojac to compute jacobians w.r.t. the last shared layer's params, and one using autogram to compute directly the gramian w.r.t. the last shared layer's params.

In GradNorm (the aggregator), we could have just one example, using torchjd.autojac.backward to accumulate jacobians in .jac, then jac_to_grad with the GradNorm aggregator, which will itself do everything.

This will not be equivalent to the examples in GradNormWeighting, because in GradNormWeighting we only consider the last shared layer's params to update the weights, instead of all params. This should be mentioned, maybe with a link to GradNormWeighting to show how to do the alternative.

"""

_initial_losses: Tensor
_initialized: Tensor
_losses: Tensor | None
_norms: Tensor | None

def __init__(self, n_tasks: int, alpha: float = 1.5) -> None:
super().__init__()
if n_tasks < 1:
raise ValueError(f"Parameter `n_tasks` must be positive. Found n_tasks={n_tasks!r}.")
self.alpha = alpha
self.weights = nn.Parameter(torch.ones(n_tasks))
self.register_buffer("_initial_losses", torch.zeros(n_tasks))
self.register_buffer("_initialized", torch.tensor(False))
self.register_buffer("_losses", None, persistent=False)
self.register_buffer("_norms", None, persistent=False)

@property
def n_tasks(self) -> int:
return self.weights.numel()

@property
def alpha(self) -> float:
return self._alpha

@alpha.setter
def alpha(self, value: float) -> None:
if not isfinite(value) or value < 0.0:
raise ValueError(f"Attribute `alpha` must be finite and non-negative. Found {value!r}.")
self._alpha = value

def set_losses(self, losses: Tensor) -> None:
"""
Stores the current unweighted task losses. The first call also records the baseline
losses. Call this before each forward, keeping the same task order throughout training.
"""
if losses.shape != self.weights.shape:
raise ValueError(f"Parameter `losses` must have shape ({self.n_tasks},).")
if losses.device != self.weights.device or losses.dtype != self.weights.dtype:
raise ValueError("Parameter `losses` must have the same device and dtype as `weights`.")
if not torch.isfinite(losses).all() or (losses < 0).any():
raise ValueError("Parameter `losses` must be finite and non-negative.")
if not self._initialized:
if (losses == 0).any():
raise ValueError("Initial losses must be strictly positive.")
self._initial_losses.copy_(losses.detach())
self._initialized.fill_(True)
self._losses = losses.detach().clone()
self._norms = None

def forward(self, gramian: PSDMatrix, /) -> Tensor:
if self._losses is None:
raise ValueError("Call `set_losses` before the forward pass.")
if gramian.shape != (self.n_tasks, self.n_tasks):
raise ValueError(
f"Parameter `gramian` must have shape ({self.n_tasks}, {self.n_tasks})."
)
if gramian.device != self.weights.device or gramian.dtype != self.weights.dtype:
raise ValueError(
"Parameter `gramian` must have the same device and dtype as `weights`."
)
self._norms = gramian.detach().diagonal().clamp_min(0).sqrt()
return self.weights.detach().clone()

def balancing_loss(self) -> Tensor:
"""
Computes the auxiliary loss using the most recent forward's gradient norms and losses.
Backpropagate this scalar before stepping the weights' optimizer. Its gradients affect
only ``weights``, even if the supplied losses or Gramian have an autograd graph.
"""
if self._losses is None or self._norms is None:
raise ValueError("Call `set_losses` and the forward pass before `balancing_loss`.")
if not self._losses.any():
return (self.weights * 0).sum()
ratios = self._losses / self._initial_losses
rates = ratios / ratios.mean()
norms = self.weights.abs() * self._norms
targets = (norms.mean() * rates.pow(self.alpha)).detach()
return (norms - targets).abs().sum()

def renormalize(self) -> None:
"""
Rescales the weights in place to sum to ``n_tasks`` after an optimizer step, as in
Algorithm 1. Raises if an update produced non-finite or negative weights, or a zero
sum. Reduce the weight learning rate if updates cross this boundary.
"""
with torch.no_grad():
total = self.weights.sum()
if not torch.isfinite(total) or (self.weights < 0).any() or total <= 0:
raise ValueError("Weights must be finite and non-negative, with a positive sum.")
self.weights.mul_(self.n_tasks / total)

def reset(self) -> None:
"""
Restores unit weights and clears the loss baseline and cached batch statistics.
Parameter identity is preserved. Reset the external optimizer separately to discard
its momentum or other state when starting a new experiment.
"""
with torch.no_grad():
self.weights.fill_(1)
self._initial_losses.zero_()
self._initialized.fill_(False)
self.weights.grad = None
self._losses = None
self._norms = None

def __repr__(self) -> str:
return f"{self.__class__.__name__}(n_tasks={self.n_tasks}, alpha={self.alpha!r})"


class GradNorm(GramianWeightedAggregator, Stateful, _NonDifferentiable):
r"""
:class:`~torchjd.Stateful` :class:`~torchjd.aggregation.GramianWeightedAggregator`
using :class:`~torchjd.aggregation.GradNormWeighting`.

Gradient norms are computed over all columns of the supplied Jacobian. For norms based on
a subset of parameters, use :class:`~torchjd.aggregation.GradNormWeighting` directly.
Pass ``parameters()`` to a separate optimizer. Call :meth:`set_losses` before aggregation,
backpropagate :meth:`balancing_loss` and call :meth:`renormalize` after the optimizer step.

:param n_tasks: Fixed positive number of tasks (rows of the Jacobian).
:param alpha: Non-negative strength of training-rate balancing.
"""

gramian_weighting: GradNormWeighting

def __init__(self, n_tasks: int, alpha: float = 1.5) -> None:
super().__init__(GradNormWeighting(n_tasks, alpha))

@property
def n_tasks(self) -> int:
return self.gramian_weighting.n_tasks

@property
def alpha(self) -> float:
return self.gramian_weighting.alpha

@alpha.setter
def alpha(self, value: float) -> None:
self.gramian_weighting.alpha = value

def set_losses(self, losses: Tensor) -> None:
"""Stores the current task losses. See :meth:`GradNormWeighting.set_losses`."""
self.gramian_weighting.set_losses(losses)

def balancing_loss(self) -> Tensor:
"""Computes the auxiliary loss. See :meth:`GradNormWeighting.balancing_loss`."""
return self.gramian_weighting.balancing_loss()

def renormalize(self) -> None:
"""Rescales the task weights. See :meth:`GradNormWeighting.renormalize`."""
self.gramian_weighting.renormalize()

def reset(self) -> None:
"""Resets the task weights and loss baseline. See :meth:`GradNormWeighting.reset`."""
self.gramian_weighting.reset()

def __repr__(self) -> str:
return f"{self.__class__.__name__}(n_tasks={self.n_tasks}, alpha={self.alpha!r})"
Loading
Loading