Skip to content

fix(aggregation): Clamp squared distances in Krum - #790

Merged
ValerianRey merged 1 commit into
SimplexLab:mainfrom
SajalDevX:fix-krum-negative-squared-distances
Oct 3, 2026
Merged

ValerianRey merged 1 commit into
SimplexLab:mainfrom
SajalDevX:fix-krum-negative-squared-distances

Conversation

@SajalDevX

Copy link
Copy Markdown
Contributor

Problem

KrumWeighting gets the pairwise squared distances from the Gramian as ||g_i||^2 + ||g_j||^2 - 2 g_i.g_j. When two rows are almost equal (which is exactly the case Krum relies on: honest gradients clustered together), rounding errors can make that value slightly negative, so torch.sqrt returns nan. torch.topk(..., largest=False) treats nan as the largest value, so the distance to the near-duplicate row gets dropped from the score of both rows. Their scores go up and Krum can pick an outlier row instead.

Repro (float32, CPU):

import torch
from torchjd.aggregation import KrumWeighting

torch.manual_seed(0)
bad = 0
for _ in range(200):
    g = torch.randn(1, 1000)
    J = torch.cat([g, g + 1e-4 * torch.randn(1, 1000), torch.randn(3, 1000)])
    w = KrumWeighting(n_byzantine=1)(J @ J.T)
    bad += int(w[0] + w[1] != 1)
print(bad)

On main this prints 43: in 43 of 200 trials one of the three random rows is selected instead of one of the two near-identical rows. With float64 inputs the same matrices always select row 0 or 1.

Solution

Clamp the squared distances to be non-negative before the square root:

distances = torch.sqrt(distances_squared.clamp(min=0.0))

With this change the repro above prints 0.

Tests

  • Added test_negative_squared_distances_are_clamped in tests/unit/aggregation/test_krum.py. It builds the Gramian of the 1-D points [1, 1, 2, 2.1, 11] and adds 1e-6 to the off-diagonal entry of the two equal points, so their squared distance is slightly negative (like a rounding error). Krum must select one of those two rows. It fails on main (weights [0, 0, 1, 0, 0]) and passes with the fix, for both float32 and PYTEST_TORCH_DTYPE=float64.
  • uv run pytest tests/unit/aggregation -W error: 1881 passed, 2 skipped (CPU only, without the cvxpy extras).
  • ruff check, ruff format --check and ty check pass on the changed files.

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

KrumWeighting computes squared distances between rows as
||g_i||^2 + ||g_j||^2 - 2 g_i.g_j from the Gramian. When two rows are
almost equal, rounding errors can make this slightly negative, so
torch.sqrt returns nan. topk treats nan as the largest value, so the
distance to the near-duplicate row is dropped from the score of both
rows and Krum can select an outlier instead.

Clamp the squared distances to be non-negative before the square root,
and add a regression test with a Gramian whose off-diagonal entry is
slightly too large.
@PierreQuinton PierreQuinton added package: aggregation cc: fix Conventional commit type for bug fixes of the actual library (changes to src). labels Oct 3, 2026

@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.

Very neat, thanks a lot. LGTM if the test used to fail.

Comment on lines +124 to +133
def test_negative_squared_distances_are_clamped() -> None:
x = tensor_([1.0, 1.0, 2.0, 2.1, 11.0])
gramian = x.unsqueeze(1) * x.unsqueeze(0)
gramian[0, 1] = gramian[0, 1] + 1e-6
gramian[1, 0] = gramian[1, 0] + 1e-6

weights = KrumWeighting(n_byzantine=1)(gramian)

assert weights[:2].sum().item() == 1.0
assert weights[2:].abs().sum().item() == 0.0

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.

Can we verify that this test indeed failed before?

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 don't have access to a computer for the weekend, so I can't double-check, but SajalDevX said it failed on main.

@ValerianRey

Copy link
Copy Markdown
Member

/opencode:Plan can you check at which commit we introduced this bug? I think we used to use cdist to compute the distances before we switched to using the gramian directly, during a refactor extracting the gramian weighting. Please also review the PR.

@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.

Very good catch, thank you! This was a nasty bug

@opencode-agent

opencode-agent Bot commented Oct 3, 2026

Copy link
Copy Markdown
Contributor

1. When was the bug introduced?

Your hypothesis is exactly right. The bug was introduced in:

  • Commit: aa87f9571c3294f4c5cf9fa06ddb69b66f557974
  • PR: Explicit dependence on gramians #280 — refactor(aggregation) Explicit dependence on gramians
  • Author/date: PierreQuinton, 2025-04-26

That commit made the dependence on the Gramian explicit in AlignedMTL, CAGrad, IMTL-G, Krum, MGDA and PCGrad. For Krum, it replaced:

distances = torch.cdist(matrix, matrix, compute_mode="donot_use_mm_for_euclid_dist")

with:

gradient_norms_squared = torch.diagonal(gramian)
distances_squared = (
    gradient_norms_squared.unsqueeze(0) + gradient_norms_squared.unsqueeze(1) - 2 * gramian
)
distances = torch.sqrt(distances_squared)

There's an irony here: the old code explicitly used compute_mode="donot_use_mm_for_euclid_dist", which exists precisely to avoid the ||x||² + ||y||² − 2⟨x,y⟩ expansion because of this exact catastrophic-cancellation issue. The refactor reintroduced that very formula by hand, without its guard. Before aa87f957, the code had been cdist-based (and safe) since the first commit of the repo, and no clamping or NaN handling ever existed — none was needed.

The unclamped sqrt then survived verbatim through the weighting-structure refactor (#347) and the rename to _krum.py (#350), dormant for ~17 months, until e885d792 (PR #790, current HEAD) added .clamp(min=0.0).

2. Review of PR #790

Verdict: LGTM — correct, minimal, well-tested. Details:

The fix is right. For a true Gramian, distances_squared ≥ 0 mathematically; negative values are pure round-off of magnitude ~eps·||g||², so clamping to 0 yields the mathematically correct distance up to round-off. This is also the standard remedy (sklearn's euclidean_distances does the same when working from the Gramian). Going back to cdist isn't an option since KrumWeighting only receives the Gramian by design — clamping is the proper fix in this architecture.

Answering the open review comment ("Can we verify that this test indeed failed before?") — I traced the test's arithmetic on the pre-fix code, deterministically (no flakiness involved):

  • d²(0,1) = 1 + 1 − 2·(1 + 1e-6). In float32, 1 + 1e-6 rounds to 1 + 8·2⁻²³, giving d² = −1.907e-6; in float64, d² = −2e-6. Negative in both dtypes → sqrt → NaN.
  • torch.topk(..., largest=False) sorts NaN as larger than any number, so the NaN mutual distance of rows 0 and 1 is never among the k smallest. Scores become [2.1, 2.1, 1.1, 1.2, 17.9] instead of [1.0, 1.0, 1.1, 1.2, 17.9], so row 2 is selected → weights [0, 0, 1, 0, 0] → both assertions fail on main, exactly as the PR description states. With the clamp, the selected row is 0 or 1 (exact tie, but tie-breaking is irrelevant since the assertions only require the winner to be among the first two rows). So the test provably fails pre-fix and passes post-fix, for both float32 and float64.

Other checks:

  • PSDMatrix is a static-typing-only marker (src/torchjd/_linalg/_matrix.py — "should never be instantiated… only used for static type checking"), so the deliberately slightly-indefinite Gramian in the test (its 2×2 principal minor has determinant −2e-6) is fine at runtime.
  • The exact-equality assertions are safe: weights come from one_hot(...).sum()/n_selected, i.e. exactly 1.0/0.0.
  • Changelog entry under [Unreleased] / Fixed is accurate and matches the project format.
  • Non-blocking observation: _imtl_g.py has a superficially similar torch.sqrt(torch.diagonal(gramian)), but it's not affected — diagonal entries are dot products of rows with themselves (sums of non-negative terms), which cannot round negative. Only the diag + diag − 2G subtraction can. No follow-up needed there.

New%20session%20-%202026-10-03T08%3A02%3A40.250Z
opencode session  |  github run

@ValerianRey
ValerianRey merged commit e142058 into SimplexLab:main Oct 3, 2026
24 of 25 checks passed
@SajalDevX

Copy link
Copy Markdown
Contributor Author

Yes, it fails on main with weights [0, 0, 1, 0, 0] (an outlier row gets selected), in both float32 and float64. Thanks for the quick review!

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