diff options
Diffstat (limited to 'tests/test_core.py')
| -rw-r--r-- | tests/test_core.py | 57 |
1 files changed, 57 insertions, 0 deletions
diff --git a/tests/test_core.py b/tests/test_core.py index 09e58f1..429eed0 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -323,3 +323,60 @@ def test_structured_phrase_encoding_separates_head_from_modifier(): np.mean([vectors[t] for t in ["red", "bus"]], axis=0), np.mean([vectors[t] for t in ["bus", "red"]], axis=0), ) + + +def test_closed_form_pair_swap_deltas_match_brute_force(): + """Every transposition's energy change, against the definition.""" + import numpy as np + from worldalign.synth_fast_gate import ( + ClosedFormEnergy, all_pair_swap_deltas, all_swaps, apply_swaps, + ) + + size = 24 + generator = torch.Generator().manual_seed(0) + def field(): + raw = torch.randn(size, size, generator=generator, dtype=torch.float64) + out = (raw + raw.T) / 2 + out.fill_diagonal_(0.0) + return out + + text, visual = field(), field() + energy = ClosedFormEnergy(text, visual, 1.0, 0.0, 64) + permutation = torch.randperm(size, generator=generator) + + base = float(energy.energy(permutation[None])[0]) + swaps = all_swaps(size, torch.device("cpu")) + brute = energy.energy(apply_swaps(permutation, swaps)) - base + + permuted = text[permutation[:, None], permutation[None, :]] + table = all_pair_swap_deltas(permuted, visual) + closed = table[swaps[:, 0], swaps[:, 1]] + + assert torch.allclose(closed, brute, atol=1e-9), ( + closed[:4].tolist(), brute[:4].tolist() + ) + + +def test_fast_pair_descent_matches_brute_force_descent(): + """The fast descent must reach the same energy as the slow one.""" + from worldalign.synth_fast_gate import ( + ClosedFormEnergy, all_swaps, fast_pair_descent, steepest_descent, + ) + + size = 24 + generator = torch.Generator().manual_seed(1) + def field(): + raw = torch.randn(size, size, generator=generator, dtype=torch.float64) + out = (raw + raw.T) / 2 + out.fill_diagonal_(0.0) + return out + + text, visual = field(), field() + energy = ClosedFormEnergy(text, visual, 1.0, 0.0, 64) + start = torch.randperm(size, generator=generator) + slow, slow_value = steepest_descent( + energy, start.clone(), all_swaps(size, torch.device("cpu")), 500 + ) + fast = fast_pair_descent(text, visual, start.clone(), 500) + fast_value = float(energy.energy(fast[None])[0]) + assert abs(fast_value - slow_value) < 1e-9, (fast_value, slow_value) |
