diff options
Diffstat (limited to 'worldalign/automorphism_probe.py')
| -rw-r--r-- | worldalign/automorphism_probe.py | 96 |
1 files changed, 96 insertions, 0 deletions
diff --git a/worldalign/automorphism_probe.py b/worldalign/automorphism_probe.py new file mode 100644 index 0000000..46949f4 --- /dev/null +++ b/worldalign/automorphism_probe.py @@ -0,0 +1,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() |
