summaryrefslogtreecommitdiff
path: root/worldalign/consensus_seed.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-01 18:53:29 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-01 18:53:29 -0500
commitde827a42e10ede662f4bd2893c4f8b6d54be45dc (patch)
treed7c2c30f5e36102572cacc70e4296d80459d992f /worldalign/consensus_seed.py
parent4f7ee05cc3b072478062e53645af016861c4b529 (diff)
Amplification beats the instance that defeated fifteen solvers
The caption-omitted field returned under 5% from every solver tried: Umeyama, GRAMPA, five Gromov-Wasserstein variants, FAQ with and without restarts, PATH convex-concave, a moment ladder, semirelaxed GW. The measurements said the answer was not a sixteenth solver. Descent on that field amplifies: a start 10% correct comes out 42%, one 20% correct comes out 78%. What no initialiser could do was clear the entry price, since all of them land in the same wrong region. So run a diverse pool of cheap descents, let them vote, round the vote matrix to a permutation by Hungarian assignment rather than argmax, descend from that, and rebuild the pool around the result. Each round feeds the amplifier a better start. 0.044 -> 0.72 on the best run, 0.54 mean over two. Competitive elsewhere: 0.977 on the easy field against 0.961 for the best GW variant. Also adds the diagnostics that led here: the basin-width probe (k=16 transpositions still returns to the exact truth 100% of the time) and the capture-threshold curve that measures amplification directly. Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/consensus_seed.py')
-rw-r--r--worldalign/consensus_seed.py180
1 files changed, 180 insertions, 0 deletions
diff --git a/worldalign/consensus_seed.py b/worldalign/consensus_seed.py
new file mode 100644
index 0000000..07a9593
--- /dev/null
+++ b/worldalign/consensus_seed.py
@@ -0,0 +1,180 @@
+"""Find a dozen pairs, not a permutation.
+
+The anchor curve changed what the problem is. Given 5% of the correspondence
+the synthetic world returns 89% of the rest, and given 10% it returns 97.6% --
+so at 256 scenes, twelve correct pairs are worth the whole search. Meanwhile
+the gate says the truth on the caption-omitted field is a strict local minimum
+that our solver misses by 0.44 in energy while landing 4.9% correct. Four point
+nine percent of 256 is twelve pairs. **The solver is already finding roughly as
+many correct pairs as a seed needs; what it cannot do is say which ones they
+are.**
+
+That is a different question and it has a cheap answer. Run many independent
+searches -- different spectral regularisers, different random starts -- and
+count how often each pairing appears. A wrong pair is wrong in a different way
+each time; a right pair is the same every time. The pairs that survive the vote
+become anchors, the rest are read off by assignment against them, and the whole
+thing is repeated with the enlarged anchor set.
+
+Seed precision is the quantity that matters, not seed recall, so it is reported
+separately at every round. Hidden pairs are read only for scoring.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+
+import numpy as np
+import torch
+from scipy.optimize import linear_sum_assignment
+
+from .common import write_json
+from .spectral_match import grampa
+from .synth_fast_gate import fast_pair_descent
+from .synth_triangle_gate import standardized
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--fields", nargs="+", required=True)
+ parser.add_argument("--labels", nargs="+", default=None)
+ parser.add_argument("--restarts", type=int, default=80)
+ parser.add_argument("--iterations", type=int, default=2000)
+ parser.add_argument("--rounds", type=int, default=4)
+ parser.add_argument("--seed-size", type=int, default=16)
+ parser.add_argument("--trials", type=int, default=3)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--output", default="artifacts/synth_v1/consensus_seed.json")
+ return parser.parse_args()
+
+
+def standardise(matrix: np.ndarray) -> np.ndarray:
+ mask = ~np.eye(len(matrix), dtype=bool)
+ values = matrix[mask]
+ out = (matrix - values.mean()) / values.std()
+ np.fill_diagonal(out, 0.0)
+ return out
+
+
+def assign_against(visual, text, anchors_left, anchors_right, probe):
+ """Match probe scenes by their profiles against an anchored coordinate frame."""
+ left = visual[np.ix_(probe, anchors_left)]
+ right = text[np.ix_(probe, anchors_right)]
+ left = left - left.mean(1, keepdims=True)
+ right = right - right.mean(1, keepdims=True)
+ left = left / left.std(1, keepdims=True).clip(1e-9)
+ right = right / right.std(1, keepdims=True).clip(1e-9)
+ similarity = left @ right.T / left.shape[1]
+ _, columns = linear_sum_assignment(-similarity)
+ margin = similarity[np.arange(len(probe)), columns] - np.partition(
+ similarity, -2, axis=1
+ )[:, -2]
+ return columns, margin
+
+
+def run_trial(visual, text, hidden, args, device) -> dict:
+ size = len(visual)
+ shuffled = text[np.ix_(hidden, hidden)]
+ shuffled_gpu = standardized(torch.from_numpy(shuffled).to(device)).double()
+ visual_gpu = standardized(torch.from_numpy(visual).to(device)).double()
+ generator = np.random.default_rng(0)
+
+ votes = np.zeros((size, size))
+ starts = [grampa(visual, shuffled, eta) for eta in (0.2, 0.5, 1.0, 2.0, 5.0)]
+ starts += [generator.permutation(size) for _ in range(args.restarts)]
+ for start in starts:
+ final = fast_pair_descent(
+ shuffled_gpu, visual_gpu,
+ torch.from_numpy(np.ascontiguousarray(start)).to(device),
+ args.iterations,
+ ).cpu().numpy()
+ votes[np.arange(size), final] += 1.0
+
+ baseline = float((hidden[votes.argmax(1)] == np.arange(size)).mean())
+
+ history = []
+ anchors_left = np.array([], dtype=int)
+ anchors_right = np.array([], dtype=int)
+ for round_index in range(args.rounds):
+ if round_index == 0:
+ confidence = votes.max(1) / len(starts)
+ proposal = votes.argmax(1)
+ else:
+ probe = np.setdiff1d(np.arange(size), anchors_left)
+ columns, margin = assign_against(
+ visual, shuffled, anchors_left, anchors_right, probe
+ )
+ confidence = np.full(size, -np.inf)
+ proposal = np.zeros(size, dtype=int)
+ confidence[probe] = margin
+ proposal[probe] = probe[columns] if False else np.array(
+ [probe[c] for c in columns]
+ )
+ # keep existing anchors
+ confidence[anchors_left] = np.inf
+ proposal[anchors_left] = anchors_right
+
+ take = min(args.seed_size * (round_index + 1), size)
+ chosen = np.argsort(confidence)[::-1][:take]
+ anchors_left = chosen
+ anchors_right = proposal[chosen]
+ correct = int((hidden[anchors_right] == anchors_left).sum())
+ precision = correct / max(len(chosen), 1)
+
+ probe = np.setdiff1d(np.arange(size), anchors_left)
+ if len(anchors_left) >= 4 and len(probe) > 1:
+ columns, _ = assign_against(
+ visual, shuffled, anchors_left, anchors_right, probe
+ )
+ resolved = np.array([probe[c] for c in columns])
+ full_correct = int((hidden[resolved] == probe).sum()) + correct
+ accuracy = full_correct / size
+ else:
+ accuracy = precision * len(chosen) / size
+ history.append({
+ "round": round_index, "anchors": int(len(chosen)),
+ "seed_precision": precision, "accuracy_after_expansion": accuracy,
+ })
+ return {"vote_baseline": baseline, "history": history}
+
+
+def main() -> None:
+ args = parse_args()
+ labels = args.labels or [p.split("/")[-1] for p in args.fields]
+ device = torch.device(args.device)
+ rows = []
+ for path, label in zip(args.fields, labels):
+ state = torch.load(path, map_location="cpu", weights_only=False)
+ visual = standardise(state["visual_field"].double().numpy())
+ text = standardise(state["text_field"].double().numpy())
+ size = len(visual)
+ print(f"\n=== {label} (N={size}, chance={1/size:.4f})", flush=True)
+ trials = []
+ for trial in range(args.trials):
+ hidden = np.random.default_rng(trial).permutation(size)
+ trials.append(run_trial(visual, text, hidden, args, device))
+ print(f" vote-only accuracy {np.mean([t['vote_baseline'] for t in trials]):.4f}",
+ flush=True)
+ for index in range(args.rounds):
+ precision = np.mean([t["history"][index]["seed_precision"] for t in trials])
+ accuracy = np.mean(
+ [t["history"][index]["accuracy_after_expansion"] for t in trials]
+ )
+ anchors = trials[0]["history"][index]["anchors"]
+ print(f" round {index}: {anchors:3d} anchors seed precision={precision:.3f}"
+ f" accuracy after expansion={accuracy:.4f}", flush=True)
+ rows.append({"field": label, "trials": trials})
+ write_json(args.output, {
+ "protocol": (
+ "Many independent descents vote on pairings; the most-agreed pairs "
+ "become anchors and the rest are assigned against them, repeated. "
+ "Hidden pairs are read only for scoring."
+ ),
+ "restarts": args.restarts, "rows": rows,
+ })
+ print(json.dumps({"done": True}))
+
+
+if __name__ == "__main__":
+ main()