summaryrefslogtreecommitdiff
path: root/worldalign/automorphism_probe.py
blob: 46949f4c576314aca13c908868814ec305fe093d (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
"""Exact automorphisms of a relation field, and the blind-recovery ceiling.

If sigma is an exact automorphism of the text field T (P_sigma T P_sigma^T = T
bit-for-bit) then the blind matching problem cannot distinguish the truth from
truth o sigma: the input (V, T) is literally unchanged.  Every functional of
(V, P T P^T) -- the pairwise energy, the third-order tr(M^3) term, any moment
in the ladder, the TLB linear term -- is invariant, so no optimiser and no
relaxation can break the tie.  The posterior over the truth is uniform on the
orbit, and the expected accuracy of *any* blind estimator is bounded by

    ceiling = (1/n) * sum_over_orbits  |orbit| * (1/|orbit|)
            = 1 - (moved - #orbits) / n

Run:  python -m worldalign.automorphism_probe artifacts/synth_v1/omit_size.pt
"""

from __future__ import annotations

import math
import sys

import numpy as np
import torch


def standardise(matrix: np.ndarray) -> np.ndarray:
    matrix = np.asarray(matrix, dtype=np.float64)
    mask = ~np.eye(len(matrix), dtype=bool)
    out = (matrix - matrix[mask].mean()) / matrix[mask].std()
    np.fill_diagonal(out, 0.0)
    return out


def automorphic_transpositions(field: np.ndarray, tol: float = 1e-9) -> np.ndarray:
    """adj[i, j] iff swapping i and j leaves the field exactly invariant.

    Rows i and j must agree on every coordinate outside {i, j}; the two
    excluded coordinates are exactly the ones the swap moves.
    """
    size = len(field)
    mismatch = np.full((size, size), np.inf)
    for i in range(size):
        gap = np.abs(field - field[i])
        gap[:, i] = 0.0
        gap[np.arange(size), np.arange(size)] = 0.0
        mismatch[i] = gap.max(axis=1)
    np.fill_diagonal(mismatch, np.inf)
    return mismatch < tol


def orbits(adjacency: np.ndarray) -> list[list[int]]:
    size = len(adjacency)
    seen = np.zeros(size, dtype=bool)
    found: list[list[int]] = []
    for start in range(size):
        if seen[start]:
            continue
        stack = [start]
        seen[start] = True
        component = [start]
        while stack:
            node = stack.pop()
            for nxt in np.nonzero(adjacency[node] & ~seen)[0]:
                seen[nxt] = True
                stack.append(nxt)
                component.append(int(nxt))
        found.append(sorted(component))
    return [c for c in found if len(c) > 1]


def report(field: np.ndarray, name: str) -> float:
    size = len(field)
    groups = orbits(automorphic_transpositions(field))
    moved = sum(len(c) for c in groups)
    log_order = sum(math.lgamma(len(c) + 1) for c in groups) / math.log(10)
    ceiling = 1.0 - (moved - len(groups)) / size
    print(
        f"{name}: {len(groups)} nontrivial orbits covering {moved}/{size} items, "
        f"sizes {sorted((len(c) for c in groups), reverse=True)[:12]}, "
        f"|Aut| >= 10^{log_order:.1f}"
    )
    print(f"{name}: blind accuracy ceiling for any energy-only method = {ceiling:.4f}")
    return ceiling


def main() -> None:
    path = sys.argv[1] if len(sys.argv) > 1 else "artifacts/synth_v1/omit_size.pt"
    state = torch.load(path, map_location="cpu", weights_only=False)
    visual = standardise(state["visual_field"].double().numpy())
    text = standardise(state["text_field"].double().numpy())
    report(text, "text field")
    report(visual, "visual field")


if __name__ == "__main__":
    main()