From e885d7925f84d78821d3160819904faafe67df75 Mon Sep 17 00:00:00 2001 From: Sajal Kumar Jana Date: Sat, 3 Oct 2026 01:34:01 +0000 Subject: [PATCH] fix(aggregation): Clamp squared distances in Krum 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. --- CHANGELOG.md | 4 ++++ src/torchjd/aggregation/_krum.py | 2 +- tests/unit/aggregation/test_krum.py | 14 +++++++++++++- 3 files changed, 18 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e25a2a708..5b0fa6cca 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -14,6 +14,10 @@ changelog does not include internal changes that do not affect the user. task had a zero gradient at the call that sets its baseline excess risk. The exponentiated gradient update is now computed in log space, so a very large excess risk saturates the weights instead of overflowing. +- Fixed `Krum` and `KrumWeighting` sometimes selecting the wrong rows when two rows of the input + 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. ## [0.17.1] - 2026-09-23 diff --git a/src/torchjd/aggregation/_krum.py b/src/torchjd/aggregation/_krum.py index 5527ab500..d12d758b6 100644 --- a/src/torchjd/aggregation/_krum.py +++ b/src/torchjd/aggregation/_krum.py @@ -29,7 +29,7 @@ def forward(self, gramian: PSDMatrix, /) -> Tensor: distances_squared = ( gradient_norms_squared.unsqueeze(0) + gradient_norms_squared.unsqueeze(1) - 2 * gramian ) - distances = torch.sqrt(distances_squared) + distances = torch.sqrt(distances_squared.clamp(min=0.0)) n_closest = gramian.shape[0] - self.n_byzantine - 2 smallest_distances, _ = torch.topk(distances, k=n_closest + 1, largest=False) diff --git a/tests/unit/aggregation/test_krum.py b/tests/unit/aggregation/test_krum.py index 6270f6778..54e775b94 100644 --- a/tests/unit/aggregation/test_krum.py +++ b/tests/unit/aggregation/test_krum.py @@ -3,7 +3,7 @@ from pytest import mark, raises from torch import Tensor from utils.contexts import ExceptionContext -from utils.tensors import ones_ +from utils.tensors import ones_, tensor_ from torchjd.aggregation import Krum from torchjd.aggregation._krum import KrumWeighting @@ -119,3 +119,15 @@ def test_weighting_n_selected_setter_rejects_non_positive() -> None: W = KrumWeighting(n_byzantine=1) with raises(ValueError, match="n_selected"): W.n_selected = 0 + + +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