Skip to content

fix(aggregation): Use Gramian dtype in AlignedMTL rank tolerance - #792

Open
SajalDevX wants to merge 1 commit into
SimplexLab:mainfrom
SajalDevX:fix-aligned-mtl-tolerance-dtype
Open

SajalDevX wants to merge 1 commit into
SimplexLab:mainfrom
SajalDevX:fix-aligned-mtl-tolerance-dtype

Conversation

@SajalDevX

Copy link
Copy Markdown
Contributor

Problem

AlignedMTLWeighting._compute_balance_transformation decides which eigenvalues of the Gramian count as non-zero with:

tol = torch.max(lambda_) * len(M) * torch.finfo().eps

torch.finfo() without an argument gives the epsilon of the default dtype (float32 in most setups), not of the Gramian. So with a float64 input, every eigenvalue below about m * 1.2e-7 times the largest one is treated as zero, even though it is perfectly representable. For tasks whose gradient norms differ by more than a factor of roughly 2000 (for m = 2), the smaller task gets dropped from the balance transformation, and the scale used by scale_mode="min" also changes. That is exactly the situation AlignedMTL is meant to handle.

Repro (CPU, default dtype float32):

import torch
from torchjd.aggregation import AlignedMTL

for ratio in [1e2, 1e3, 1e4]:
    J = torch.tensor([[1.0, 0.0], [0.0, 1.0 / ratio]], dtype=torch.float64)
    print(ratio, AlignedMTL()(J).tolist())

On main:

100.0 [0.005, 0.005]
1000.0 [0.0005, 0.0005]
10000.0 [0.5, 0.0]

The last line should be [5e-05, 5e-05]. Instead, the second task is ignored completely and the result is half the first gradient.

Solution

Use the epsilon of the Gramian's dtype:

tol = torch.max(lambda_) * len(M) * torch.finfo(M.dtype).eps

For float32 inputs with the default dtype left at float32, the tolerance stays exactly the same. With this change the repro prints [5e-05, 5e-05] for the last case.

Tests

  • Added test_float64_small_eigenvalue_is_kept in tests/unit/aggregation/test_aligned_mtl.py. It aggregates the float64 matrix [[1, 0], [0, 1e-4]] and checks that the result is [5e-5, 5e-5]. It fails on main (it returns [0.5, 0]) and passes with the fix. The matrix is built in float64 on purpose, so the test checks the same thing under float32 and under PYTEST_TORCH_DTYPE=float64.
  • pytest tests/unit/aggregation -W error: 1507 passed, 15 skipped, for both float32 and PYTEST_TORCH_DTYPE=float64 (CPU only, without the optional cvxpy/qpsolvers extras).
  • ruff check, ruff format --check and ty check pass on the changed files.

Added an entry under [Unreleased] / Fixed in the changelog.

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.

@PierreQuinton PierreQuinton left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Good catch, LGTM. Just to be sure /opencode:Plan Can you check if this is applicable anywhere else in the aggregation package?

@PierreQuinton PierreQuinton added package: aggregation cc: fix Conventional commit type for bug fixes of the actual library (changes to src). labels Oct 4, 2026
@ValerianRey

Copy link
Copy Markdown
Member

/opencode:Plan please review this. Has the bug always been there or has it been introduced in a non-initial commit? If so, which one?

@ValerianRey ValerianRey left a comment

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.

Good catch thanks! Small comment and we can merge

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

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

@opencode-agent

opencode-agent Bot commented Oct 4, 2026

Copy link
Copy Markdown
Contributor

Verified both key facts directly against GitHub. Here's the answer:

The bug has always been there — since the very first commit

The torch.finfo() call without a dtype argument was not introduced by any later commit. It exists verbatim in the repository's root commit:

  • Commit: 611105b7744cb663a63fc883417ac5ab0d879156 — "parents": [] confirms it's the root
  • Author: Valérian Rey (co-author: Pierre Quinton)
  • Date: 2024-06-01
  • Message: "Add everything" — this commit squashed 1600+ private development commits, so no earlier public history exists

The file (then src/torchjd/aggregation/aligned_mtl.py, class AlignedMTLWrapper) already contained the exact buggy line in that commit:

tol = torch.max(lambda_) * len(M) * torch.finfo().eps

I confirmed this by fetching the file content at that commit — the line is byte-identical to what main had before this PR.

Additional context

  • Since the root commit is a squash, we can't pinpoint which private commit introduced it, but given the file header ("Partly adapted from SamsungLabs/MTL — MIT License, Copyright (c) 2022 Samsung"), it was almost certainly inherited from the original Aligned-MTL reference implementation rather than written fresh in TorchJD.
  • The line survived ~2.3 years and ~28 commits touching the file completely untouched — including the scale_mode feature refactor in 8653929a (feat(aggregation) Add scale_mode parameter to AlignedMTL #527, 2026-01-23), which rewrote the surrounding code (replacing lambda_R = lambda_[-1] with the scale branching) but left the tolerance line as-is. The only change ever made to it is this PR's fix (e5c8e58d).

So: no non-initial commit to blame — the fix in this PR is the first time this tolerance computation has ever been touched.

New%20session%20-%202026-10-04T14%3A27%3A10.055Z
opencode session  |  github run

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

cc: fix Conventional commit type for bug fixes of the actual library (changes to src). package: aggregation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants