summaryrefslogtreecommitdiff
path: root/tests/test_core.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_core.py')
-rw-r--r--tests/test_core.py57
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)