Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 4 additions & 3 deletions src/torchjd/aggregation/_excess_mtl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
43 changes: 43 additions & 0 deletions tests/unit/aggregation/test_excess_mtl.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from typing import cast

import torch
from pytest import mark, raises
from torch import Tensor
Expand Down Expand Up @@ -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)
Loading