feat(aggregation): Add GradNorm - #785
giovannicozzolongo wants to merge 2 commits into
Conversation
Signed-off-by: Giovanni Cozzolongo <79092266+giovannicozzolongo@users.noreply.github.com>
Signed-off-by: Giovanni Cozzolongo <79092266+giovannicozzolongo@users.noreply.github.com>
|
@giovannicozzolongo thanks a lot for the PR! I start reviewing it now. It's a quite hard topic so I may take a few days. Thanks for your patience! |
|
/opencode:Plan please review this PR. In particular, discuss its API, and compare it with the other possible API where the weight optimizer would be internal, and the usage would look something like: losses = ...
gramian = ...
gradnorm_weighting.set_losses(losses)
weights = gradnorm_weighting(gramian)
# computes aux loss, updates weights using internal optimizer, and returns previous weights
losses.backward(weights)Make a table with the pros and cons of both options. If you think about other options, mention them. |
ToggleI've now studied the PR thoroughly: the new Here is my review, focused on the API. 1. What the PR's API actually isThe PR proposes an external-optimizer, split-phase protocol. Per iteration, the user writes: weight_optimizer.zero_grad()
weighting.set_losses(losses) # side channel: current unweighted losses
weights = weighting(gramian) # caches norms; returns current weights (detached clone)
weighting.balancing_loss().backward() # populates weights.grad only
losses.backward(weights) # model backward with the *same* (pre-update) weights
model_optimizer.step()
weight_optimizer.step()
weighting.renormalize() # restore sum(w) = n_tasksThe alternative under discussion (internal optimizer) collapses the weight update into the forward call: weighting.set_losses(losses)
weights = weighting(gramian) # aux loss, backward, internal step, renormalize; returns pre-update weights
losses.backward(weights)2. A constraint that shapes both designs
A second observation, also inherent to both designs: GradNorm's weights are learned state, not a per-step function of the Gramian. The forward is therefore degenerate ("return current weights") and the Gramian only influences future calls. The real computation lives in 3. Comparison table
4. The decisive factors
5. Other options worth considering
6. RecommendationKeep Option A (the PR's design). Its costs are ergonomic and recoverable (Option C can be added non-breakingly later if users ask); Option B's costs are structural (checkpointing, AMP/scheduler friction, hidden side effects in 7. Minor, non-blocking observations
Great work overall — the API protocol is unusual for this library, but it's unusual because GradNorm itself doesn't fit the "weights = f(Gramian)" mold, and the chosen design makes the right trade-offs explicit. |
|
Please disregard opencode's comment unless I confirm its concerns in the review. Also note that opencode is wrong about there being no precedent of having an optimizer inside a nn.Module in TorchJD. We have exactly that in the FAMO scalarizer. IMO option B is much closer to TorchJD's usual interface, but comes with 4 minor drawbacks:
If 3) and 4) can be solved I would go for option B. /opencode:Plan please investigate about 3) and 4). |
ToggleBoth concerns are investigable from the codebase itself — FAMO ( TL;DR
Concern 3 —
|
| Drawback | Solution | Cost |
|---|---|---|
3) _NonDifferentiable |
FAMO's enable_grad + autograd.grad island |
Plus one genuinely new but standard piece: a self.training gate for eval/inference-mode |
4) state_dict |
Tier 1 override (~15 lines), or Tier 0 = FAMO status quo | One docstring caveat about non-tensor state dict values |
Neither is structural. If the maintainer preference is Option B provided 3) and 4) are solvable, I'd say the condition is met — with the caveat that the training-mode gate should be called out explicitly in the docs, since no other TorchJD aggregator currently changes behavior under .eval().
ValerianRey
left a comment
There was a problem hiding this comment.
Based on my discussion with OpenCode, I'd like to change the API to something where the weight optimizer is internal only.
For the NonDifferentiableMixin issue, a torch.enable_grad() island, like in FAMO, seems correct. I think having a .eval() and .train() method, with the same interface as those present in some pytorch modules, would also be helpful.
For the StateDict issue, just having an optimizer property to let the user get the optimizer themselves, and save its state dict manually, seems like the best option.
Could you try making the changes to reach this new interface?
| | [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) | |
There was a problem hiding this comment.
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).
There was a problem hiding this comment.
We usually have method-specific examples directly in the docstring of the class of the method. Could you move it there?
| model_optimizer.zero_grad() | ||
| weight_optimizer.zero_grad() | ||
| losses = model(features).square().mean(dim=0) | ||
| jacs = jac(losses, list(model.parameters()), retain_graph=True) |
There was a problem hiding this comment.
To match the paper a bit better, we could compute only the jacobians w.r.t. the last shared layer here.
| See :doc:`the GradNorm example <../../examples/gradnorm>` for computing the norms | ||
| using only the last shared layer while weighting all model parameters. |
There was a problem hiding this comment.
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.


Refs #665 (GradNorm).
Adds
GradNormWeightingand itsGradNormaggregator wrapper. Task weights arenn.Parameters initialised to one.set_lossesrecords the current losses, the forward caches detached gradient norms and returns detached weights, andbalancing_lossprovides the auxiliary objective for an external optimizer.renormalizerestores the weights' sum after each optimizer step.This follows Algorithm 1's direct weight updates. The balancing target is held fixed during differentiation, and the auxiliary backward cannot change model gradients. The initial loss baseline is saved in the state dict. The number of tasks stays fixed so optimizer references remain valid across calls and resets.
The example computes norms on the last shared layer with
autogram.Engineand applies the learned weights to the whole model. No new dependencies are required.The optimizer placement and API remain open for review. This draft uses an external optimizer and direct weights rather than LibMTL's softmax parameterisation.
Testing
ty check,ruff checkandruff format --check: passed.Tensorwithout importing it on Python 3.12; the existing Lightning example emits a model summary where no output is expected.git diff --check: passed.Tested with Python 3.12.13 and PyTorch 2.14.0+cu130. The minimum supported Python and PyTorch versions were not tested.