fix(aggregation): Clamp squared distances in Krum - #790
Conversation
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
left a comment
There was a problem hiding this comment.
Very neat, thanks a lot. LGTM if the test used to fail.
| 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 |
There was a problem hiding this comment.
Can we verify that this test indeed failed before?
There was a problem hiding this comment.
I don't have access to a computer for the weekend, so I can't double-check, but SajalDevX said it failed on main.
|
/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
left a comment
There was a problem hiding this comment.
Very good catch, thank you! This was a nasty bug
1. When was the bug introduced?Your hypothesis is exactly right. The bug was introduced in:
That commit made the dependence on the Gramian explicit in 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 The unclamped 2. Review of PR #790Verdict: LGTM — correct, minimal, well-tested. Details: The fix is right. For a true Gramian, 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):
Other checks:
|
|
Yes, it fails on main with weights |

Problem
KrumWeightinggets 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, sotorch.sqrtreturnsnan.torch.topk(..., largest=False)treatsnanas 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):
On
mainthis 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:
With this change the repro above prints 0.
Tests
test_negative_squared_distances_are_clampedintests/unit/aggregation/test_krum.py. It builds the Gramian of the 1-D points[1, 1, 2, 2.1, 11]and adds1e-6to 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 onmain(weights[0, 0, 1, 0, 0]) and passes with the fix, for bothfloat32andPYTEST_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 --checkandty checkpass on the changed files.Added an entry under
[Unreleased] / Fixedin the changelog.