Conversation
If a task has a zero gradient at the call that sets its baseline excess risk, every later call divides that task's excess risk by ~1e-7, exp() overflows to inf, and the normalisation turns all weights into nan for the rest of the run. Compute the exponentiated gradient update in log space (softmax of log-weights plus the step) so that a huge excess risk saturates the weights instead of overflowing. On well-behaved inputs the result is unchanged.
SajalDevX
requested review from
a team,
KhusPatel4450,
PierreQuinton and
ValerianRey
as code owners
September 29, 2026 13:15
Member
ValerianRey
approved these changes
Sep 29, 2026
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
If a task has a zero gradient at the call that sets its baseline excess risk (a head that is still frozen, a loss that is not active yet, an all-zero first batch),
ExcessMTLWeightingnever recovers:_initial_w[1]is0, so from the next call onw[1] / (0 + 1e-7)is ~1e7,torch.exp(w * eta)overflows toinf, andweights / weights.sum()isinf / inf = nan. The same happens for an all-zero matrix at the first call.Fix
Do the exponentiated gradient update in log space:
This is the same update (
weights * exp(eta * w), normalised) whenever the old one was finite —test_log_space_update_matches_direct_formulachecks that against the direct formula — but a huge excess risk now saturates the weights (that task → 1, the others → 0) instead of overflowing. It does not change the baseline handling, so behaviour stays aligned with the official implementation / LibMTL apart from no longer producingnan.Tests
test_zero_baseline_excess_risk_keeps_weights_finiteandtest_all_zero_matrix_at_baseline_keeps_weights_finite: both fail onmainwithnan.test_log_space_update_matches_direct_formula: equivalence on regular inputs.tests/unit/aggregation: 963 passed. Changelog entry added under Unreleased / Fixed.