diff --git a/CHANGELOG.md b/CHANGELOG.md index e25a2a70..5b0fa6cc 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 5527ab50..d12d758b 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 6270f677..54e775b9 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