diff --git a/CHANGELOG.md b/CHANGELOG.md index 69116c24..bcc353e0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/README.md b/README.md index ff862bc4..af85f96c 100644 --- a/README.md +++ b/README.md @@ -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) | | [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) | diff --git a/docs/source/docs/aggregation/gradnorm.rst b/docs/source/docs/aggregation/gradnorm.rst new file mode 100644 index 00000000..144f0a34 --- /dev/null +++ b/docs/source/docs/aggregation/gradnorm.rst @@ -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 diff --git a/docs/source/docs/aggregation/index.rst b/docs/source/docs/aggregation/index.rst index b77332db..11cf695e 100644 --- a/docs/source/docs/aggregation/index.rst +++ b/docs/source/docs/aggregation/index.rst @@ -33,6 +33,7 @@ Abstract base classes excess_mtl.rst fairgrad.rst graddrop.rst + gradnorm.rst gradvac.rst imtl_g.rst krum.rst diff --git a/docs/source/examples/gradnorm.rst b/docs/source/examples/gradnorm.rst new file mode 100644 index 00000000..aca5df26 --- /dev/null +++ b/docs/source/examples/gradnorm.rst @@ -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. diff --git a/docs/source/examples/index.rst b/docs/source/examples/index.rst index 603c8a3f..466cd326 100644 --- a/docs/source/examples/index.rst +++ b/docs/source/examples/index.rst @@ -18,6 +18,8 @@ This section contains some usage examples for TorchJD. - :doc:`Multi-Task Learning (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 ` 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) ` 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>`. @@ -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 diff --git a/src/torchjd/aggregation/__init__.py b/src/torchjd/aggregation/__init__.py index 55e855ee..16a7bca0 100644 --- a/src/torchjd/aggregation/__init__.py +++ b/src/torchjd/aggregation/__init__.py @@ -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 @@ -80,6 +81,8 @@ "FairGrad", "FairGradWeighting", "GradDrop", + "GradNorm", + "GradNormWeighting", "GradVac", "GradVacWeighting", "GramianWeightedAggregator", diff --git a/src/torchjd/aggregation/_gradnorm.py b/src/torchjd/aggregation/_gradnorm.py new file mode 100644 index 00000000..3773e628 --- /dev/null +++ b/src/torchjd/aggregation/_gradnorm.py @@ -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 + `_ (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) + 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. + """ + + _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})" diff --git a/tests/unit/aggregation/test_gradnorm.py b/tests/unit/aggregation/test_gradnorm.py new file mode 100644 index 00000000..56ea70f2 --- /dev/null +++ b/tests/unit/aggregation/test_gradnorm.py @@ -0,0 +1,336 @@ +from copy import deepcopy + +import torch +from pytest import mark, raises +from settings import DEVICE, DTYPE +from torch import Tensor, nn +from torch.optim import SGD, Adam +from torch.testing import assert_close +from utils.tensors import eye_, ones_, randn_, tensor_, zeros_ + +from torchjd.aggregation import GradNorm, GradNormWeighting +from torchjd.autogram import Engine +from torchjd.autojac import jac + +from ._asserts import assert_expected_structure, assert_non_differentiable +from ._inputs import scaled_matrices, typical_matrices + + +def test_representations() -> None: + assert repr(GradNormWeighting(3)) == "GradNormWeighting(n_tasks=3, alpha=1.5)" + assert repr(GradNorm(3, alpha=0.0)) == "GradNorm(n_tasks=3, alpha=0.0)" + assert str(GradNorm(3)) == "GradNorm" + + +@mark.parametrize("matrix", typical_matrices + scaled_matrices) +def test_expected_structure(matrix: Tensor) -> None: + aggregator = GradNorm(matrix.shape[0]).to(device=DEVICE, dtype=DTYPE) + aggregator.set_losses(ones_(matrix.shape[0])) + assert_expected_structure(aggregator, matrix) + + +def test_initial_weights_and_detached_target() -> None: + weighting = GradNormWeighting(3).to(device=DEVICE, dtype=DTYPE) + losses = tensor_([2.0, 4.0, 8.0], requires_grad=True) + gramian = torch.diag(tensor_([1.0, 4.0, 25.0])).requires_grad_() + weighting.set_losses(losses) + assert_close(weighting(gramian), ones_(3)) + balancing_loss = weighting.balancing_loss() + assert_close(balancing_loss, tensor_(14 / 3)) + balancing_loss.backward() + assert_close(weighting.weights.grad, tensor_([-1.0, -2.0, 5.0])) + assert losses.grad is None + assert gramian.grad is None + + +@mark.parametrize("alpha", [0.0, 1.0, 1.5]) +def test_training_rates_use_initial_losses(alpha: float) -> None: + weighting = GradNormWeighting(3, alpha=alpha).to(device=DEVICE, dtype=DTYPE) + initial = tensor_([2.0, 4.0, 8.0]) + weighting.set_losses(initial) + initial.mul_(10) + weighting.set_losses(tensor_([2.0, 2.0, 2.0])) + weighting(torch.diag(tensor_([1.0, 4.0, 25.0]))) + targets = (8 / 3) * tensor_([12 / 7, 6 / 7, 3 / 7]).pow(alpha) + expected = (tensor_([1.0, 2.0, 5.0]) - targets).abs().sum() + assert_close(weighting.balancing_loss(), expected) + + +def test_two_sgd_steps_match_algorithm_one() -> None: + weighting = GradNormWeighting(3, alpha=0.0).to(device=DEVICE, dtype=DTYPE) + parameter = weighting.weights + optimizer = SGD(weighting.parameters(), lr=0.1) + diagonals = [tensor_([1.0, 4.0, 25.0]), tensor_([4.0, 1.0, 16.0])] + expected = [tensor_([33 / 28, 9 / 7, 15 / 28]), tensor_([411 / 350, 291 / 175, 57 / 350])] + for diagonal, weights in zip(diagonals, expected, strict=True): + optimizer.zero_grad() + weighting.set_losses(ones_(3)) + before = weighting(torch.diag(diagonal)) + weighting.balancing_loss().backward() + optimizer.step() + weighting.renormalize() + assert_close(weighting.weights, weights) + assert_close(weighting.weights.sum(), tensor_(3.0)) + assert (weighting.weights >= 0).all() + assert not before.requires_grad + assert not torch.allclose(before, weighting.weights) + assert weighting.weights is parameter + + +def test_auxiliary_backward_does_not_change_model_gradients() -> None: + parameter = nn.Parameter(tensor_([1.0, 2.0])) + losses = torch.stack([parameter.square().sum(), 3 * (parameter - 1).square().sum()]) + weighting = GradNormWeighting(2).to(device=DEVICE, dtype=DTYPE) + with torch.no_grad(): + weighting.weights.copy_(tensor_([0.5, 1.5])) + J = torch.stack([torch.autograd.grad(loss, parameter, retain_graph=True)[0] for loss in losses]) + weighting.set_losses(losses) + weights = weighting(J @ J.T) + losses.backward(weights) + expected_grad = tensor_([1.0, 11.0]) + assert_close(parameter.grad, expected_grad) + assert weighting.weights.grad is None + weighting.balancing_loss().backward() + assert_close(parameter.grad, expected_grad) + assert weighting.weights.grad is not None + + +def test_last_shared_layer_matches_direct_autograd() -> None: + shared = nn.Sequential(nn.Linear(3, 4), nn.Tanh(), nn.Linear(4, 2)).to(DEVICE, DTYPE) + heads = nn.ModuleList([nn.Linear(2, 1), nn.Linear(2, 1)]).to(DEVICE, DTYPE) + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + optimizer = Adam(weighting.parameters(), lr=0.001) + reference_weights = nn.Parameter(ones_(2)) + reference_optimizer = Adam([reference_weights], lr=0.001) + parameters = [*shared.parameters(), *heads.parameters()] + engine = Engine(shared[2], batch_dim=None) + initial_losses = None + for _ in range(3): + optimizer.zero_grad() + reference_optimizer.zero_grad() + for parameter in parameters: + parameter.grad = None + representation = shared(randn_(4, 3)) + losses = torch.stack([head(representation).square().mean() for head in heads]) + if initial_losses is None: + initial_losses = losses.detach().clone() + direct_norms = [] + for weight, loss in zip(reference_weights, losses, strict=True): + gradients = torch.autograd.grad( + weight * loss, list(shared[2].parameters()), retain_graph=True, create_graph=True + ) + direct_norms.append(torch.cat([g.flatten() for g in gradients]).norm()) + norms = torch.stack(direct_norms) + ratios = losses.detach() / initial_losses + targets = (norms.mean() * (ratios / ratios.mean()).pow(1.5)).detach() + reference_loss = (norms - targets).abs().sum() + reference_gradient = torch.autograd.grad(reference_loss, reference_weights)[0] + expected_model = torch.autograd.grad( + losses, parameters, grad_outputs=reference_weights.detach(), retain_graph=True + ) + weighting.set_losses(losses) + weights = weighting(engine.compute_gramian(losses)) + assert_close(weighting.balancing_loss(), reference_loss) + weighting.balancing_loss().backward() + assert_close(weighting.weights.grad, reference_gradient) + assert all(parameter.grad is None for parameter in parameters) + losses.backward(weights) + for parameter, expected in zip(parameters, expected_model, strict=True): + assert_close(parameter.grad, expected) + optimizer.step() + weighting.renormalize() + reference_weights.grad = reference_gradient + reference_optimizer.step() + with torch.no_grad(): + reference_weights.mul_(2 / reference_weights.sum()) + assert_close(weighting.weights, reference_weights) + + +def test_aggregator_matches_weighting_and_updates() -> None: + aggregator = GradNorm(3).to(DEVICE, DTYPE) + weighting = GradNormWeighting(3).to(DEVICE, DTYPE) + optimizers = [SGD(module.parameters(), lr=0.01) for module in (aggregator, weighting)] + for _ in range(2): + J = randn_(3, 4) + losses = ones_(3) + aggregator.set_losses(losses) + weighting.set_losses(losses) + assert_close(aggregator(J), weighting(J @ J.T) @ J) + assert_close(aggregator.balancing_loss(), weighting.balancing_loss()) + for module, optimizer in zip((aggregator, weighting), optimizers, strict=True): + optimizer.zero_grad() + module.balancing_loss().backward() + optimizer.step() + module.renormalize() + aggregator.reset() + aggregator.set_losses(losses) + assert_close(aggregator(J), J.sum(dim=0)) + + +def test_non_differentiable() -> None: + aggregator = GradNorm(3).to(DEVICE, DTYPE) + aggregator.set_losses(ones_(3)) + assert_non_differentiable(aggregator, ones_(3, 5, requires_grad=True)) + + +@mark.parametrize("n_columns", [0, 4]) +def test_zero_gradients(n_columns: int) -> None: + aggregator = GradNorm(2).to(DEVICE, DTYPE) + aggregator.set_losses(ones_(2)) + assert_close(aggregator(zeros_(2, n_columns)), zeros_(n_columns)) + loss = aggregator.balancing_loss() + assert_close(loss, tensor_(0.0)) + loss.backward() + assert_close(aggregator.gramian_weighting.weights.grad, zeros_(2)) + + +def test_single_task() -> None: + weighting = GradNormWeighting(1).to(DEVICE, DTYPE) + weighting.set_losses(ones_(1)) + assert_close(weighting(eye_(1)), ones_(1)) + weighting.balancing_loss().backward() + assert_close(weighting.weights.grad, zeros_(1)) + + +@mark.parametrize( + "losses, expected_gradient", [([0.0, 1.0], [1.0, -1.0]), ([0.0, 0.0], [0.0, 0.0])] +) +def test_zero_current_losses(losses: list[float], expected_gradient: list[float]) -> None: + weighting = GradNormWeighting(2, alpha=1.0).to(DEVICE, DTYPE) + weighting.set_losses(ones_(2)) + weighting.set_losses(tensor_(losses)) + weighting(eye_(2)) + weighting.balancing_loss().backward() + assert_close(weighting.weights.grad, tensor_(expected_gradient)) + + +def test_checkpoint_restores_weights_and_baseline() -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + optimizer = Adam(weighting.parameters(), lr=0.01) + weighting.set_losses(tensor_([2.0, 3.0])) + weighting(torch.diag(tensor_([1.0, 9.0]))) + weighting.balancing_loss().backward() + optimizer.step() + weighting.renormalize() + restored = GradNormWeighting(2).to(DEVICE, DTYPE) + restored.load_state_dict(deepcopy(weighting.state_dict())) + restored_optimizer = Adam(restored.parameters(), lr=0.01) + restored_optimizer.load_state_dict(deepcopy(optimizer.state_dict())) + with raises(ValueError, match="set_losses"): + restored(eye_(2)) + for module, opt in ((weighting, optimizer), (restored, restored_optimizer)): + opt.zero_grad() + module.set_losses(tensor_([1.0, 1.0])) + module(torch.diag(tensor_([4.0, 1.0]))) + module.balancing_loss().backward() + opt.step() + module.renormalize() + assert_close(weighting.weights, restored.weights) + + +def test_reset_keeps_parameter_and_replaces_baseline() -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + parameter = weighting.weights + optimizer = SGD(weighting.parameters(), lr=0.1) + weighting.set_losses(tensor_([2.0, 3.0])) + weighting(torch.diag(tensor_([1.0, 9.0]))) + weighting.balancing_loss().backward() + optimizer.step() + weighting.reset() + assert weighting.weights is parameter + assert weighting.weights.grad is None + assert optimizer.param_groups[0]["params"][0] is parameter + fresh = GradNormWeighting(2).to(DEVICE, DTYPE) + for module in (weighting, fresh): + module.set_losses(tensor_([5.0, 1.0])) + module(torch.diag(tensor_([1.0, 9.0]))) + assert_close(weighting.weights, fresh.weights) + assert_close(weighting.balancing_loss(), fresh.balancing_loss()) + + +def test_to_moves_baseline_and_batch_statistics() -> None: + weighting = GradNormWeighting(2).to(device=DEVICE, dtype=torch.float32) + weighting.set_losses(tensor_([2.0, 3.0]).float()) + weighting(eye_(2).float()) + weighting = weighting.double() + assert weighting.balancing_loss().dtype == torch.float64 + weighting.set_losses(tensor_([1.0, 1.0]).double()) + assert weighting(eye_(2).double()).dtype == torch.float64 + + +@mark.parametrize("cls", [GradNorm, GradNormWeighting]) +@mark.parametrize("n_tasks", [0, -1]) +def test_invalid_task_count(cls: type[GradNorm | GradNormWeighting], n_tasks: int) -> None: + with raises(ValueError, match="n_tasks"): + cls(n_tasks) + + +@mark.parametrize("cls", [GradNorm, GradNormWeighting]) +@mark.parametrize("alpha", [-1.0, float("nan"), float("inf")]) +def test_alpha_validation(cls: type[GradNorm | GradNormWeighting], alpha: float) -> None: + with raises(ValueError, match="alpha"): + cls(2, alpha=alpha) + module = cls(2) + module.alpha = 0.5 + assert module.alpha == 0.5 + with raises(ValueError, match="alpha"): + module.alpha = alpha + + +@mark.parametrize("losses", [[0.0, 1.0], [-1.0, 1.0], [float("nan"), 1.0], [float("inf"), 1.0]]) +def test_invalid_initial_losses(losses: list[float]) -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + with raises(ValueError, match="losses"): + weighting.set_losses(tensor_(losses)) + + +def test_shape_and_call_order_validation() -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + with raises(ValueError, match="set_losses"): + weighting(eye_(2)) + with raises(ValueError, match="forward"): + weighting.balancing_loss() + with raises(ValueError, match="shape"): + weighting.set_losses(ones_(3)) + weighting.set_losses(ones_(2)) + with raises(ValueError, match="shape"): + weighting(eye_(3)) + weighting(eye_(2)) + weighting.set_losses(ones_(2)) + with raises(ValueError, match="forward"): + weighting.balancing_loss() + + +def test_dtype_validation() -> None: + weighting = GradNormWeighting(2).to(device=DEVICE, dtype=torch.float64) + with raises(ValueError, match="dtype"): + weighting.set_losses(ones_(2).float()) + weighting.set_losses(ones_(2).double()) + with raises(ValueError, match="dtype"): + weighting(eye_(2).float()) + + +@mark.parametrize("weights", [[0.0, 0.0], [-1.0, 2.0], [float("nan"), 1.0], [float("inf"), 1.0]]) +def test_invalid_weight_updates(weights: list[float]) -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + with torch.no_grad(): + weighting.weights.copy_(tensor_(weights)) + with raises(ValueError, match="Weights"): + weighting.renormalize() + + +def test_fresh_checkpoint_can_initialize() -> None: + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + weighting.load_state_dict(GradNormWeighting(2).state_dict()) + weighting.set_losses(ones_(2)) + assert_close(weighting(eye_(2)), ones_(2)) + + +def test_autojac_usage() -> None: + parameter = nn.Parameter(tensor_([1.0, 2.0])) + losses = parameter.square() + J = jac(losses, [parameter], retain_graph=True)[0] + weighting = GradNormWeighting(2).to(DEVICE, DTYPE) + weighting.set_losses(losses) + losses.backward(weighting(J @ J.T)) + assert_close(parameter.grad, tensor_([2.0, 4.0]))