summaryrefslogtreecommitdiff
path: root/artifacts/spectral_frontier_probe/eval_impure_anchors.py
diff options
context:
space:
mode:
Diffstat (limited to 'artifacts/spectral_frontier_probe/eval_impure_anchors.py')
-rw-r--r--artifacts/spectral_frontier_probe/eval_impure_anchors.py65
1 files changed, 65 insertions, 0 deletions
diff --git a/artifacts/spectral_frontier_probe/eval_impure_anchors.py b/artifacts/spectral_frontier_probe/eval_impure_anchors.py
new file mode 100644
index 0000000..f6b4397
--- /dev/null
+++ b/artifacts/spectral_frontier_probe/eval_impure_anchors.py
@@ -0,0 +1,65 @@
+import numpy as np, torch
+from scipy.optimize import linear_sum_assignment
+
+def standardise(M):
+ M = np.asarray(M, dtype=np.float64); mask = ~np.eye(len(M), dtype=bool); v = M[mask]
+ out = (M - v.mean()) / v.std(); np.fill_diagonal(out, 0.0); return out
+
+d = torch.load('/home/yurenh2/emm/artifacts/synth_v1/omit_size.pt', map_location='cpu', weights_only=False)
+V = standardise(d['visual_field']); T = standardise(d['text_field']); N = len(V)
+truth = np.arange(N)
+
+def expand(anchorL, anchorR):
+ """anchorL: vision indices; anchorR: claimed text partner. Returns full accuracy."""
+ probe = np.setdiff1d(np.arange(N), anchorL)
+ # probe scenes on the text side are whatever is not claimed
+ probeR = np.setdiff1d(np.arange(N), anchorR)
+ L = V[np.ix_(probe, anchorL)]; R = T[np.ix_(probeR, anchorR)]
+ L = (L - L.mean(1, keepdims=True)) / L.std(1, keepdims=True).clip(1e-9)
+ R = (R - R.mean(1, keepdims=True)) / R.std(1, keepdims=True).clip(1e-9)
+ S = L @ R.T / L.shape[1]
+ _, cols = linear_sum_assignment(-S)
+ resolved = probeR[cols]
+ n_ok = int((resolved == probe).sum()) + int((anchorR == anchorL).sum())
+ return n_ok / N
+
+rng = np.random.default_rng(0)
+REP = 20
+print("=== (a) PURE random anchors: reproduce the published curve ===")
+for K in (12, 25, 40, 50, 64, 128):
+ a = [expand(*(lambda s: (s, s))(np.sort(rng.choice(N, K, replace=False)))) for _ in range(REP)]
+ print(f" K={K:4d} pure accuracy {np.mean(a):.3f} +- {np.std(a):.3f}")
+
+print("\n=== (b) IMPURE random anchors: c correct, K-c wrong (wrong = derangement among selected) ===")
+for K in (12, 25, 40, 50, 100):
+ for prec in (1.0, 0.92, 0.88, 0.84, 0.80, 0.72, 0.60):
+ nb = int(round(K * (1 - prec)))
+ if nb == 1: nb = 2 # a single wrong pair is impossible inside a bijection
+ accs = []
+ for _ in range(REP):
+ sel = np.sort(rng.choice(N, K, replace=False))
+ right = sel.copy()
+ if nb >= 2:
+ bad = rng.choice(K, nb, replace=False)
+ sh = right[bad].copy()
+ while True:
+ perm = rng.permutation(nb)
+ if not (perm == np.arange(nb)).any(): break
+ right[bad] = sh[perm]
+ accs.append(expand(sel, right))
+ print(f" K={K:4d} prec={prec:.2f} ({K-nb}/{K} right) accuracy {np.mean(accs):.3f} +- {np.std(accs):.3f}")
+ print()
+
+print("=== (c) IMPURE with wrong partners drawn from OUTSIDE the anchor set ===")
+for K in (25, 50):
+ for prec in (1.0, 0.88, 0.84, 0.80, 0.72):
+ nb = int(round(K * (1 - prec))); accs = []
+ for _ in range(REP):
+ sel = np.sort(rng.choice(N, K, replace=False))
+ right = sel.copy()
+ if nb:
+ bad = rng.choice(K, nb, replace=False)
+ outside = np.setdiff1d(np.arange(N), sel)
+ right[bad] = rng.choice(outside, nb, replace=False)
+ accs.append(expand(sel, right))
+ print(f" K={K:4d} prec={prec:.2f} accuracy {np.mean(accs):.3f} +- {np.std(accs):.3f}")