Skip to content

fix(aggregation): Apply ExcessMTL exponentiation in log domain - #787

Open
SajalDevX wants to merge 2 commits into
SimplexLab:mainfrom
SajalDevX:fix/excess-mtl-nan-weights
Open

SajalDevX wants to merge 2 commits into
SimplexLab:mainfrom
SajalDevX:fix/excess-mtl-nan-weights

Conversation

@SajalDevX

Copy link
Copy Markdown

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), ExcessMTLWeighting never recovers:

>>> a = ExcessMTL()
>>> J = torch.randn(2, 5); J[1] = 0
>>> a(J)
tensor([ 0.5143, -1.3016,  0.1904,  0.5383, -2.2004])
>>> a(torch.randn(2, 5))
tensor([nan, nan, nan, nan, nan])       # and every call after this

_initial_w[1] is 0, so from the next call on w[1] / (0 + 1e-7) is ~1e7, torch.exp(w * eta) overflows to inf, and weights / weights.sum() is inf / inf = nan. The same happens for an all-zero matrix at the first call.

Fix

Do the exponentiated gradient update in log space:

weights = torch.softmax(torch.log(weights) + w * self._robust_step_size, dim=0)

This is the same update (weights * exp(eta * w), normalised) whenever the old one was finite — test_log_space_update_matches_direct_formula checks 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 producing nan.

Tests

  • test_zero_baseline_excess_risk_keeps_weights_finite and test_all_zero_matrix_at_baseline_keeps_weights_finite: both fail on main with nan.
  • test_log_space_update_matches_direct_formula: equivalence on regular inputs.
  • tests/unit/aggregation: 963 passed. Changelog entry added under Unreleased / Fixed.

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.
@ValerianRey ValerianRey added package: aggregation cc: fix Conventional commit type for bug fixes of the actual library (changes to src). labels Sep 29, 2026
@ValerianRey ValerianRey changed the title Keep ExcessMTL weights finite when a baseline excess risk is zero Apply ExcessMTL exponentiation in log domain Sep 29, 2026
@github-actions github-actions Bot changed the title Apply ExcessMTL exponentiation in log domain fix(aggregation): Apply ExcessMTL exponentiation in log domain Sep 29, 2026
@ValerianRey

Copy link
Copy Markdown
Member

cc @KhusPatel4450

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

cc: fix Conventional commit type for bug fixes of the actual library (changes to src). package: aggregation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants