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
4 changes: 4 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion src/torchjd/aggregation/_aligned_mtl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
9 changes: 8 additions & 1 deletion tests/unit/aggregation/test_aligned_mtl.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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))
Comment on lines +65 to +68

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think we should make this test agnostic of the dtype (i.e. rename it test_smal_eigenvalue_is_kept, and not use dtype=torch.float64). One of our CI runs uses dtype float64, so it will be tested on both float32 and float64

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks! I tried that, but a dtype-agnostic version doesn't catch the bug: it only shows when the matrix dtype differs from torch's default dtype, because torch.finfo() falls back to the default. With the default dtype, the old and new tolerances are identical, so the test would also pass on main (in the float64 CI run the default is float64 too). That's why the matrix is explicitly float64: under the float32 run it reproduces the bug, and under the float64 run it is just a normal case. I could rename it to test_small_eigenvalue_is_kept_for_non_default_dtype to make that clearer. Would that work for you?

@ValerianRey ValerianRey Oct 6, 2026 •

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

When our CI runs with PYTEST_TORCH_DTYPE=float64, it doesn't change torch's default dtype. It just makes tensor_ (and many other functions, defined in tests/utils/tensors.py) implicitly use dtype=float64. See tests/settings.py. So I don't think the float64 CI run is supposed to pass on main. Are you sure it does pass?

Loading