|
5 | 5 | import numpy as np |
6 | 6 | import pepsy as py |
7 | 7 | import pytest |
8 | | -from pathlib import Path |
9 | | -import runpy |
| 8 | +import warnings |
10 | 9 |
|
11 | 10 | qtn = pytest.importorskip("quimb.tensor") |
| 11 | +from quimb.tensor.belief_propagation import D1BP # noqa: E402 |
12 | 12 |
|
13 | 13 | from pepsy.bp import RelayBPResult, one_norm_bp, relay_bp # noqa: E402 |
14 | 14 | from pepsy.bp.relay import _relay_message_sources # noqa: E402 |
@@ -43,6 +43,26 @@ def _peps_2x2(): |
43 | 43 | return qtn.PEPS.rand(Lx=2, Ly=2, bond_dim=2, seed=10) |
44 | 44 |
|
45 | 45 |
|
| 46 | +def _odd_antiferromagnetic_cycle(epsilon): |
| 47 | + """Return a positive triangle whose edge factors nearly flip a bit.""" |
| 48 | + factor = np.array([[epsilon, 1.0], [1.0, epsilon]]) |
| 49 | + return qtn.TensorNetwork( |
| 50 | + [ |
| 51 | + qtn.Tensor(factor, inds=("ab", "ca")), |
| 52 | + qtn.Tensor(factor, inds=("ab", "bc")), |
| 53 | + qtn.Tensor(factor, inds=("bc", "ca")), |
| 54 | + ] |
| 55 | + ) |
| 56 | + |
| 57 | + |
| 58 | +def _polarized_messages(tn): |
| 59 | + """Choose a deterministic non-fixed initial message for every edge end.""" |
| 60 | + return { |
| 61 | + key: np.array([1.0, 0.0]) |
| 62 | + for key in D1BP(tn, update="parallel").messages |
| 63 | + } |
| 64 | + |
| 65 | + |
46 | 66 | def test_one_norm_bp_close_to_exact_on_small_grid(): |
47 | 67 | tn = _ising_tn(3, 0.2) |
48 | 68 | exact = float(tn.contract()) |
@@ -232,13 +252,48 @@ def test_parallel_update_runs(): |
232 | 252 |
|
233 | 253 | def test_relay_d1bp_odd_cycle_stress_cases_converge_strictly(): |
234 | 254 | """Relay damps deterministic parallel D1BP stalls on odd parity cycles.""" |
235 | | - example_path = ( |
236 | | - Path(__file__).resolve().parents[1] |
237 | | - / "examples" |
238 | | - / "RelayBP" |
239 | | - / "odd_cycle_stress.py" |
240 | | - ) |
241 | | - records = runpy.run_path(str(example_path))["run_stress_cases"]() |
| 255 | + records = [] |
| 256 | + for epsilon in (1e-3, 1e-2): |
| 257 | + tn = _odd_antiferromagnetic_cycle(epsilon) |
| 258 | + exact = float(tn.contract()) |
| 259 | + initial = _polarized_messages(tn) |
| 260 | + common = { |
| 261 | + "method": "d1bp", |
| 262 | + "init_messages": initial, |
| 263 | + "update": "parallel", |
| 264 | + "diis": False, |
| 265 | + "max_iterations": 100, |
| 266 | + "tol": 1e-10, |
| 267 | + } |
| 268 | + with warnings.catch_warnings(): |
| 269 | + warnings.filterwarnings( |
| 270 | + "ignore", |
| 271 | + message="Belief propagation did not converge.*", |
| 272 | + category=UserWarning, |
| 273 | + ) |
| 274 | + plain = one_norm_bp(tn, **common) |
| 275 | + relay = relay_bp( |
| 276 | + tn, |
| 277 | + **common, |
| 278 | + num_relays=5, |
| 279 | + memory_first_leg=True, |
| 280 | + gamma_range=(0.2, 0.8), |
| 281 | + seed=0, |
| 282 | + ) |
| 283 | + relay_estimate = float(relay.contract()) |
| 284 | + records.append( |
| 285 | + { |
| 286 | + "epsilon": epsilon, |
| 287 | + "exact": exact, |
| 288 | + "plain_converged": plain.converged, |
| 289 | + "plain_max_mdiff": plain.max_mdiff, |
| 290 | + "relay_converged": relay.converged, |
| 291 | + "relay_iterations": relay.iterations, |
| 292 | + "relay_num_legs": relay.num_legs_run, |
| 293 | + "relay_max_mdiff": relay.max_mdiff, |
| 294 | + "relay_relative_error": abs(relay_estimate - exact) / abs(exact), |
| 295 | + } |
| 296 | + ) |
242 | 297 |
|
243 | 298 | assert len(records) == 2 |
244 | 299 | for record in records: |
|
0 commit comments