From e5c8e58d51635b362b21c5fda8646a84f2351782 Mon Sep 17 00:00:00 2001 From: Sajal Kumar Jana Date: Sun, 4 Oct 2026 04:03:09 +0000 Subject: [PATCH] fix(aggregation): Use Gramian dtype in AlignedMTL rank tolerance AlignedMTLWeighting computed the eigenvalue cutoff used to find the rank of the Gramian with torch.finfo().eps, i.e. the machine epsilon of the default dtype (float32 in most setups), whatever the dtype of the Gramian. With a float64 input, eigenvalues smaller than about m * 1.2e-7 times the largest one were treated as zero, so a task whose gradient is a few thousand times smaller than the others was dropped from the balance transformation entirely. Use torch.finfo(M.dtype).eps instead. --- CHANGELOG.md | 4 ++++ src/torchjd/aggregation/_aligned_mtl.py | 2 +- tests/unit/aggregation/test_aligned_mtl.py | 9 ++++++++- 3 files changed, 13 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5b0fa6cca..961d35cd7 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. ## [0.17.1] - 2026-09-23 diff --git a/src/torchjd/aggregation/_aligned_mtl.py b/src/torchjd/aggregation/_aligned_mtl.py index da994d853..fec6b2000 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 6eacfba9b..3608a3eab 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))