From 402bbb9a0af4f9d7b16b94249a54cf019ff111b3 Mon Sep 17 00:00:00 2001 From: Sajal Kumar Jana Date: Tue, 29 Sep 2026 13:15:18 +0000 Subject: [PATCH 1/2] Keep ExcessMTL weights finite when a baseline excess risk is zero 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. --- CHANGELOG.md | 7 ++++ src/torchjd/aggregation/_excess_mtl.py | 7 ++-- tests/unit/aggregation/test_excess_mtl.py | 41 +++++++++++++++++++++++ 3 files changed, 52 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 69116c24..e25a2a70 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,13 @@ changelog does not include internal changes that do not affect the user. ## [Unreleased] +### Fixed + +- Fixed `ExcessMTL` and `ExcessMTLWeighting` producing `nan` weights for the rest of the run when a + task had a zero gradient at the call that sets its baseline excess risk. The exponentiated + gradient update is now computed in log space, so a very large excess risk saturates the weights + instead of overflowing. + ## [0.17.1] - 2026-09-23 ### Fixed diff --git a/src/torchjd/aggregation/_excess_mtl.py b/src/torchjd/aggregation/_excess_mtl.py index ff8062fd..40b8a807 100644 --- a/src/torchjd/aggregation/_excess_mtl.py +++ b/src/torchjd/aggregation/_excess_mtl.py @@ -120,10 +120,11 @@ def forward(self, matrix: Matrix, /) -> Tensor: else: w = w / (self._initial_w + 1e-7) # Scale processing (Section 3.2) - # Exponentiated gradient weight update (Equation 9) + # Exponentiated gradient weight update (Equation 9), done in log space so that a very + # large excess risk (e.g. a task whose baseline excess risk was zero) saturates the + # weights instead of overflowing to inf and turning them into nan for good. weights = cast(Tensor, self._weights) - weights = weights * torch.exp(w * self._robust_step_size) - weights = weights / weights.sum() + weights = torch.softmax(torch.log(weights) + w * self._robust_step_size, dim=0) self._weights = weights return weights diff --git a/tests/unit/aggregation/test_excess_mtl.py b/tests/unit/aggregation/test_excess_mtl.py index 04367062..095c21f6 100644 --- a/tests/unit/aggregation/test_excess_mtl.py +++ b/tests/unit/aggregation/test_excess_mtl.py @@ -233,3 +233,44 @@ def test_excess_mtl_reset_delegates() -> None: agg(J) agg.reset() assert_close(first, agg(J)) + + +def test_zero_baseline_excess_risk_keeps_weights_finite() -> None: + """A task with a zero gradient at the baseline call must not poison the weights with nan.""" + weighting = ExcessMTLWeighting() + first = randn_(2, 5) + first[1] = 0.0 + weighting(first) + + weights = weighting(randn_(2, 5)) + + assert torch.isfinite(weights).all() + assert_close(weights.sum(), tensor_(1.0)) + assert weights[1] > weights[0] + + +def test_all_zero_matrix_at_baseline_keeps_weights_finite() -> None: + weighting = ExcessMTLWeighting() + weighting(torch.zeros(3, 4)) + + weights = weighting(randn_(3, 4)) + + assert torch.isfinite(weights).all() + assert_close(weights.sum(), tensor_(1.0)) + + +def test_log_space_update_matches_direct_formula() -> None: + """On well-behaved inputs the update is the same as weights * exp(eta * w), normalized.""" + weighting = ExcessMTLWeighting(robust_step_size=0.5) + weighting(randn_(3, 6)) + before = weighting._weights.clone() + matrix = randn_(3, 6) + + weights = weighting(matrix) + + sq_grad_sum = weighting._sq_grad_sum + w = (matrix**2 / torch.sqrt(sq_grad_sum + 1e-7)).sum(dim=1) + w = w / (weighting._initial_w + 1e-7) + expected = before * torch.exp(0.5 * w) + expected = expected / expected.sum() + assert_close(weights, expected) From 09b56d1d91784baab770975f2a7a9ad36291835a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Val=C3=A9rian=20Rey?= Date: Tue, 29 Sep 2026 16:35:35 +0200 Subject: [PATCH 2/2] Fix linting errors --- tests/unit/aggregation/test_excess_mtl.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/tests/unit/aggregation/test_excess_mtl.py b/tests/unit/aggregation/test_excess_mtl.py index 095c21f6..c46c5a3f 100644 --- a/tests/unit/aggregation/test_excess_mtl.py +++ b/tests/unit/aggregation/test_excess_mtl.py @@ -1,3 +1,5 @@ +from typing import cast + import torch from pytest import mark, raises from torch import Tensor @@ -263,14 +265,14 @@ def test_log_space_update_matches_direct_formula() -> None: """On well-behaved inputs the update is the same as weights * exp(eta * w), normalized.""" weighting = ExcessMTLWeighting(robust_step_size=0.5) weighting(randn_(3, 6)) - before = weighting._weights.clone() + before = cast(Tensor, weighting._weights).clone() matrix = randn_(3, 6) weights = weighting(matrix) - sq_grad_sum = weighting._sq_grad_sum + sq_grad_sum = cast(Tensor, weighting._sq_grad_sum) w = (matrix**2 / torch.sqrt(sq_grad_sum + 1e-7)).sum(dim=1) - w = w / (weighting._initial_w + 1e-7) + w = w / (cast(Tensor, weighting._initial_w) + 1e-7) expected = before * torch.exp(0.5 * w) expected = expected / expected.sum() assert_close(weights, expected)