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..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 @@ -233,3 +235,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 = cast(Tensor, weighting._weights).clone() + matrix = randn_(3, 6) + + weights = weighting(matrix) + + sq_grad_sum = cast(Tensor, weighting._sq_grad_sum) + w = (matrix**2 / torch.sqrt(sq_grad_sum + 1e-7)).sum(dim=1) + w = w / (cast(Tensor, weighting._initial_w) + 1e-7) + expected = before * torch.exp(0.5 * w) + expected = expected / expected.sum() + assert_close(weights, expected)