-
Notifications
You must be signed in to change notification settings - Fork 24
feat(aggregation): Add GradNorm #785
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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 |
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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?
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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. |
| 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) | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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})" | ||
There was a problem hiding this comment.
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).