diff --git a/CHANGELOG.md b/CHANGELOG.md index 170f0530..413d3c4f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,6 +18,10 @@ changelog does not include internal changes that do not affect the user. matrix are almost equal. Rounding errors could make the squared distance between such rows slightly negative, giving a `nan` distance that was then ignored when computing the scores. Squared distances are now clamped to be non-negative before taking the square root. +- Fixed `AlignedMTL` and `AlignedMTLWeighting` ignoring tasks with a much smaller gradient than + the others when the input is in `float64`. The tolerance used to find the rank of the Gramian was + always based on the machine epsilon of the default dtype (usually `float32`) instead of the dtype + of the Gramian, so valid small eigenvalues were discarded. - Fixed `Engine.compute_gramian` silently ignoring the contribution of a module when one of its forward passes had no associated backward pass (e.g. because its output was detached). It now raises a `ValueError` instead. diff --git a/src/torchjd/aggregation/_aligned_mtl.py b/src/torchjd/aggregation/_aligned_mtl.py index 5cb389c0..7df1558c 100644 --- a/src/torchjd/aggregation/_aligned_mtl.py +++ b/src/torchjd/aggregation/_aligned_mtl.py @@ -61,7 +61,7 @@ def _compute_balance_transformation( scale_mode: SUPPORTED_SCALE_MODE = "min", ) -> Tensor: lambda_, V = torch.linalg.eigh(M, UPLO="U") # More modern equivalent to torch.symeig - tol = torch.max(lambda_) * len(M) * torch.finfo().eps + tol = torch.max(lambda_) * len(M) * torch.finfo(M.dtype).eps rank = sum(lambda_ > tol) if rank == 0: diff --git a/tests/unit/aggregation/test_aligned_mtl.py b/tests/unit/aggregation/test_aligned_mtl.py index 6eacfba9..3608a3ea 100644 --- a/tests/unit/aggregation/test_aligned_mtl.py +++ b/tests/unit/aggregation/test_aligned_mtl.py @@ -1,7 +1,8 @@ import torch from pytest import mark, raises from torch import Tensor -from utils.tensors import ones_ +from torch.testing import assert_close +from utils.tensors import ones_, tensor_ from torchjd.aggregation import AlignedMTL, ConstantWeighting @@ -59,3 +60,9 @@ def test_scale_mode_setter_updates_value() -> None: A.scale_mode = "rmse" assert A.scale_mode == "rmse" assert A.gramian_weighting.scale_mode == "rmse" + + +def test_float64_small_eigenvalue_is_kept() -> None: + J = tensor_([[1.0, 0.0], [0.0, 1e-4]], dtype=torch.float64) + result = AlignedMTL()(J) + assert_close(result, tensor_([5e-5, 5e-5], dtype=torch.float64))