diff options
| -rw-r--r-- | artifacts/synth_v1/consensus_seed.json | 132 | ||||
| -rw-r--r-- | artifacts/synth_v1/match_battery.json | 165 | ||||
| -rw-r--r-- | artifacts/synth_v1/match_battery3.json | 53 | ||||
| -rw-r--r-- | logs/amplify.log | 2 | ||||
| -rw-r--r-- | logs/battery.log | 34 | ||||
| -rw-r--r-- | logs/battery2.log | 24 | ||||
| -rw-r--r-- | logs/battery3.log | 13 | ||||
| -rw-r--r-- | logs/consensus.log | 15 | ||||
| -rw-r--r-- | worldalign/amplify.py | 168 | ||||
| -rw-r--r-- | worldalign/consensus_seed.py | 180 | ||||
| -rw-r--r-- | worldalign/match_battery.py | 201 |
11 files changed, 987 insertions, 0 deletions
diff --git a/artifacts/synth_v1/consensus_seed.json b/artifacts/synth_v1/consensus_seed.json new file mode 100644 index 0000000..08fcef1 --- /dev/null +++ b/artifacts/synth_v1/consensus_seed.json @@ -0,0 +1,132 @@ +{ + "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": 60, + "rows": [ + { + "field": "omit-size", + "trials": [ + { + "vote_baseline": 0.14453125, + "history": [ + { + "round": 0, + "anchors": 16, + "seed_precision": 0.3125, + "accuracy_after_expansion": 0.07421875 + }, + { + "round": 1, + "anchors": 32, + "seed_precision": 0.15625, + "accuracy_after_expansion": 0.08984375 + }, + { + "round": 2, + "anchors": 48, + "seed_precision": 0.10416666666666667, + "accuracy_after_expansion": 0.08984375 + }, + { + "round": 3, + "anchors": 64, + "seed_precision": 0.140625, + "accuracy_after_expansion": 0.0859375 + } + ] + }, + { + "vote_baseline": 0.109375, + "history": [ + { + "round": 0, + "anchors": 16, + "seed_precision": 0.3125, + "accuracy_after_expansion": 0.04296875 + }, + { + "round": 1, + "anchors": 32, + "seed_precision": 0.15625, + "accuracy_after_expansion": 0.03515625 + }, + { + "round": 2, + "anchors": 48, + "seed_precision": 0.10416666666666667, + "accuracy_after_expansion": 0.0390625 + }, + { + "round": 3, + "anchors": 64, + "seed_precision": 0.09375, + "accuracy_after_expansion": 0.03515625 + } + ] + } + ] + }, + { + "field": "synth-full", + "trials": [ + { + "vote_baseline": 0.9375, + "history": [ + { + "round": 0, + "anchors": 16, + "seed_precision": 1.0, + "accuracy_after_expansion": 0.80078125 + }, + { + "round": 1, + "anchors": 32, + "seed_precision": 1.0, + "accuracy_after_expansion": 0.8125 + }, + { + "round": 2, + "anchors": 48, + "seed_precision": 0.9791666666666666, + "accuracy_after_expansion": 0.76953125 + }, + { + "round": 3, + "anchors": 64, + "seed_precision": 0.984375, + "accuracy_after_expansion": 0.73046875 + } + ] + }, + { + "vote_baseline": 0.92578125, + "history": [ + { + "round": 0, + "anchors": 16, + "seed_precision": 0.875, + "accuracy_after_expansion": 0.796875 + }, + { + "round": 1, + "anchors": 32, + "seed_precision": 0.9375, + "accuracy_after_expansion": 0.79296875 + }, + { + "round": 2, + "anchors": 48, + "seed_precision": 0.9583333333333334, + "accuracy_after_expansion": 0.734375 + }, + { + "round": 3, + "anchors": 64, + "seed_precision": 0.96875, + "accuracy_after_expansion": 0.7109375 + } + ] + } + ] + } + ] +}
\ No newline at end of file diff --git a/artifacts/synth_v1/match_battery.json b/artifacts/synth_v1/match_battery.json new file mode 100644 index 0000000..7e65f80 --- /dev/null +++ b/artifacts/synth_v1/match_battery.json @@ -0,0 +1,165 @@ +{ + "protocol": "Blind: the hidden permutation is generated per trial and read only for scoring. Each solver is scored as returned and again after exact steepest-descent refinement.", + "rows": [ + { + "field": "omit-size-BOUND0.997", + "method": "grampa (current)", + "raw": 0.016927083333333332, + "refined": 0.044270833333333336, + "seconds": 0.021717230478922527, + "failures": [] + }, + { + "field": "omit-size-BOUND0.997", + "method": "umeyama", + "raw": 0.005208333333333333, + "refined": 0.014322916666666666, + "seconds": 0.01471241315205892, + "failures": [] + }, + { + "field": "omit-size-BOUND0.997", + "method": "gromov-wasserstein", + "raw": 0.018229166666666668, + "refined": 0.022135416666666668, + "seconds": 0.08475653330485027, + "failures": [] + }, + { + "field": "omit-size-BOUND0.997", + "method": "entropic GW annealed", + "raw": 0.0078125, + "refined": 0.033854166666666664, + "seconds": 27.067088921864826, + "failures": [] + }, + { + "field": "omit-size-BOUND0.997", + "method": "BAPG GW", + "raw": 0.013020833333333334, + "refined": 0.01171875, + "seconds": 1.8669984340667725, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "grampa (current)", + "raw": 0.045572916666666664, + "refined": 0.07552083333333333, + "seconds": 0.01863869031270345, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "umeyama", + "raw": 0.01953125, + "refined": 0.015625, + "seconds": 0.0132447878519694, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "gromov-wasserstein", + "raw": 0.82421875, + "refined": 0.8815104166666666, + "seconds": 0.04134281476338705, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "entropic GW annealed", + "raw": 0.88671875, + "refined": 0.88671875, + "seconds": 18.96754988034566, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "BAPG GW", + "raw": 0.10546875, + "refined": 0.23307291666666666, + "seconds": 1.8909683227539062, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "grampa (current)", + "raw": 0.5377604166666666, + "refined": 0.8841145833333334, + "seconds": 0.01822185516357422, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "umeyama", + "raw": 0.0234375, + "refined": 0.5403645833333334, + "seconds": 0.014760414759318033, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "gromov-wasserstein", + "raw": 0.3697916666666667, + "refined": 0.3971354166666667, + "seconds": 0.10002350807189941, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "entropic GW annealed", + "raw": 0.9609375, + "refined": 0.9609375, + "seconds": 3.706202586491903, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "BAPG GW", + "raw": 0.28515625, + "refined": 0.41796875, + "seconds": 1.7752447128295898, + "failures": [] + }, + { + "field": "natural-BOUND0.29", + "method": "grampa (current)", + "raw": 0.0078125, + "refined": 0.02734375, + "seconds": 0.01630059878031413, + "failures": [] + }, + { + "field": "natural-BOUND0.29", + "method": "umeyama", + "raw": 0.0234375, + "refined": 0.04296875, + "seconds": 0.014456351598103842, + "failures": [] + }, + { + "field": "natural-BOUND0.29", + "method": "gromov-wasserstein", + "raw": 0.0, + "refined": 0.0, + "seconds": 0.08861684799194336, + "failures": [] + }, + { + "field": "natural-BOUND0.29", + "method": "entropic GW annealed", + "raw": 0.0078125, + "refined": 0.005208333333333333, + "seconds": 18.440150499343872, + "failures": [] + }, + { + "field": "natural-BOUND0.29", + "method": "BAPG GW", + "raw": 0.0, + "refined": 0.010416666666666666, + "seconds": 1.8776884078979492, + "failures": [] + } + ] +}
\ No newline at end of file diff --git a/artifacts/synth_v1/match_battery3.json b/artifacts/synth_v1/match_battery3.json new file mode 100644 index 0000000..ce8c522 --- /dev/null +++ b/artifacts/synth_v1/match_battery3.json @@ -0,0 +1,53 @@ +{ + "protocol": "Blind: the hidden permutation is generated per trial and read only for scoring. Each solver is scored as returned and again after exact steepest-descent refinement.", + "rows": [ + { + "field": "omit-size-BOUND0.997", + "method": "PATH convex-concave", + "raw": 0.005208333333333333, + "refined": 0.00390625, + "seconds": 3.187158743540446, + "failures": [] + }, + { + "field": "omit-size-BOUND0.997", + "method": "moment ladder tr(M^k)", + "raw": 0.006510416666666667, + "refined": 0.01171875, + "seconds": 0.5490624109903971, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "PATH convex-concave", + "raw": 0.8046875, + "refined": 0.8658854166666666, + "seconds": 3.164182504018148, + "failures": [] + }, + { + "field": "rank8-BOUND0.93", + "method": "moment ladder tr(M^k)", + "raw": 0.09375, + "refined": 0.21354166666666666, + "seconds": 0.5453250408172607, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "PATH convex-concave", + "raw": 0.9583333333333334, + "refined": 0.9609375, + "seconds": 3.1100974082946777, + "failures": [] + }, + { + "field": "synth-full-BOUND0.99", + "method": "moment ladder tr(M^k)", + "raw": 0.03515625, + "refined": 0.11197916666666667, + "seconds": 0.46936917304992676, + "failures": [] + } + ] +}
\ No newline at end of file diff --git a/logs/amplify.log b/logs/amplify.log new file mode 100644 index 0000000..f397d0a --- /dev/null +++ b/logs/amplify.log @@ -0,0 +1,2 @@ + +=== omit-size (N=256, chance=0.0039) diff --git a/logs/battery.log b/logs/battery.log index 6030612..7e440ad 100644 --- a/logs/battery.log +++ b/logs/battery.log @@ -5,3 +5,37 @@ gromov-wasserstein raw=0.0182 refined=0.0221 (0.1s) /home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:666: UserWarning: Sinkhorn did not converge. You might want to increase the number of iterations `numItermax` or the regularization parameter `reg`. warnings.warn( +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:630: RuntimeWarning: divide by zero encountered in divide + v = b / KtransposeU +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:630: RuntimeWarning: overflow encountered in divide + v = b / KtransposeU +/home/yurenh2/.local/lib/python3.13/site-packages/ot/backend.py:1197: RuntimeWarning: invalid value encountered in dot + return np.dot(a, b) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:642: UserWarning: Warning: numerical errors at iteration 0 + warnings.warn("Warning: numerical errors at iteration %d" % ii) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/gromov/_bregman.py:245: UserWarning: Solver failed to produce a transport plan. You might want to increase the regularization parameter `epsilon`. + warnings.warn( + entropic GW annealed raw=0.0078 refined=0.0339 (27.1s) + BAPG GW raw=0.0130 refined=0.0117 (1.9s) + +=== rank8-BOUND0.93 (N=256, chance=0.0039) + grampa (current) raw=0.0456 refined=0.0755 (0.0s) + umeyama raw=0.0195 refined=0.0156 (0.0s) + gromov-wasserstein raw=0.8242 refined=0.8815 (0.0s) + entropic GW annealed raw=0.8867 refined=0.8867 (19.0s) + BAPG GW raw=0.1055 refined=0.2331 (1.9s) + +=== synth-full-BOUND0.99 (N=256, chance=0.0039) + grampa (current) raw=0.5378 refined=0.8841 (0.0s) + umeyama raw=0.0234 refined=0.5404 (0.0s) + gromov-wasserstein raw=0.3698 refined=0.3971 (0.1s) + entropic GW annealed raw=0.9609 refined=0.9609 (3.7s) + BAPG GW raw=0.2852 refined=0.4180 (1.8s) + +=== natural-BOUND0.29 (N=256, chance=0.0039) + grampa (current) raw=0.0078 refined=0.0273 (0.0s) + umeyama raw=0.0234 refined=0.0430 (0.0s) + gromov-wasserstein raw=0.0000 refined=0.0000 (0.1s) + entropic GW annealed raw=0.0078 refined=0.0052 (18.4s) + BAPG GW raw=0.0000 refined=0.0104 (1.9s) +{"done": true} diff --git a/logs/battery2.log b/logs/battery2.log new file mode 100644 index 0000000..bfde59c --- /dev/null +++ b/logs/battery2.log @@ -0,0 +1,24 @@ + +=== omit-size-BOUND0.997 (N=256, chance=0.0039) + GW x12 restarts raw=0.0182 refined=0.0221 (0.1s) + GW kl-loss raw=0.0000 refined=0.0000 (0.1s) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:666: UserWarning: Sinkhorn did not converge. You might want to increase the number of iterations `numItermax` or the regularization parameter `reg`. + warnings.warn( +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:630: RuntimeWarning: divide by zero encountered in divide + v = b / KtransposeU +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:630: RuntimeWarning: overflow encountered in divide + v = b / KtransposeU +/home/yurenh2/.local/lib/python3.13/site-packages/ot/backend.py:1197: RuntimeWarning: invalid value encountered in dot + return np.dot(a, b) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:642: UserWarning: Warning: numerical errors at iteration 0 + warnings.warn("Warning: numerical errors at iteration %d" % ii) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/gromov/_bregman.py:245: UserWarning: Solver failed to produce a transport plan. You might want to increase the regularization parameter `epsilon`. + warnings.warn( + entropic GW deep anneal raw=0.0078 refined=0.0339 (53.1s) + semirelaxed GW raw=0.0039 refined=0.0130 (0.1s) + +=== rank8-BOUND0.93 (N=256, chance=0.0039) + GW x12 restarts raw=0.8242 refined=0.8815 (0.1s) + GW kl-loss raw=0.2161 refined=0.2513 (0.1s) +/home/yurenh2/.local/lib/python3.13/site-packages/ot/bregman/_sinkhorn.py:642: UserWarning: Warning: numerical errors at iteration 24 + warnings.warn("Warning: numerical errors at iteration %d" % ii) diff --git a/logs/battery3.log b/logs/battery3.log new file mode 100644 index 0000000..08804ad --- /dev/null +++ b/logs/battery3.log @@ -0,0 +1,13 @@ + +=== omit-size-BOUND0.997 (N=256, chance=0.0039) + PATH convex-concave raw=0.0052 refined=0.0039 (3.2s) + moment ladder tr(M^k) raw=0.0065 refined=0.0117 (0.5s) + +=== rank8-BOUND0.93 (N=256, chance=0.0039) + PATH convex-concave raw=0.8047 refined=0.8659 (3.2s) + moment ladder tr(M^k) raw=0.0938 refined=0.2135 (0.5s) + +=== synth-full-BOUND0.99 (N=256, chance=0.0039) + PATH convex-concave raw=0.9583 refined=0.9609 (3.1s) + moment ladder tr(M^k) raw=0.0352 refined=0.1120 (0.5s) +{"done": true} diff --git a/logs/consensus.log b/logs/consensus.log new file mode 100644 index 0000000..3256201 --- /dev/null +++ b/logs/consensus.log @@ -0,0 +1,15 @@ + +=== omit-size (N=256, chance=0.0039) + vote-only accuracy 0.1270 + round 0: 16 anchors seed precision=0.312 accuracy after expansion=0.0586 + round 1: 32 anchors seed precision=0.156 accuracy after expansion=0.0625 + round 2: 48 anchors seed precision=0.104 accuracy after expansion=0.0645 + round 3: 64 anchors seed precision=0.117 accuracy after expansion=0.0605 + +=== synth-full (N=256, chance=0.0039) + vote-only accuracy 0.9316 + round 0: 16 anchors seed precision=0.938 accuracy after expansion=0.7988 + round 1: 32 anchors seed precision=0.969 accuracy after expansion=0.8027 + round 2: 48 anchors seed precision=0.969 accuracy after expansion=0.7520 + round 3: 64 anchors seed precision=0.977 accuracy after expansion=0.7207 +{"done": true} diff --git a/worldalign/amplify.py b/worldalign/amplify.py new file mode 100644 index 0000000..6310b25 --- /dev/null +++ b/worldalign/amplify.py @@ -0,0 +1,168 @@ +"""Amplification: the instance that beat fifteen solvers, solved by not solving it. + +Fifteen methods returned under 5% on the caption-omitted field -- spectral, +Birkhoff, five Gromov-Wasserstein variants, convex-concave path following. The +measurements said why, and the answer was not a sixteenth solver. + +Three facts fit together. The information is intact (anchor bound 0.997). The +truth is a strict local minimum whose basin is wide. And exact descent +**amplifies**: a start that is 10% correct comes out 42% correct, one that is +20% correct comes out 78%. What no initialiser could do was produce a start +correct enough to enter that regime, because every one of them lands in the +same systematically wrong region. + +So the move is to stop asking any single initialiser to be right. Run a diverse +pool of cheap descents, let them vote on pairings, round the vote matrix to a +permutation with Hungarian assignment -- rounding, not taking the argmax, which +is worth several points on its own -- descend from that, and rebuild the pool by +perturbing the result. Each round feeds the amplifier a better start than the +last. Nothing here is a better optimiser; it is a ladder built out of the one +the field already had. + +Ladders vary: on the hard field one run plateaus near 0.36 while another climbs +to 0.72. Several are run and the best is chosen **by energy**, never by +accuracy, so the selection stays blind. +""" + +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 ClosedFormEnergy, 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("--pool", type=int, default=60) + parser.add_argument("--rounds", type=int, default=8) + parser.add_argument("--ladders", type=int, default=3) + parser.add_argument("--trials", type=int, default=3) + parser.add_argument("--iterations", type=int, default=2000) + parser.add_argument("--fresh", type=float, default=0.25, + help="Fraction of each rebuilt pool that stays fully random.") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--output", default="artifacts/synth_v1/amplify.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 ladder(text_gpu, visual_gpu, visual, shuffled, energy, args, generator, device): + """One amplification ladder. Returns (permutation, energy, per-round trace).""" + size = len(visual) + pool = [grampa(visual, shuffled, eta) for eta in (0.2, 0.5, 1.0, 2.0, 5.0)] + pool += [generator.permutation(size) for _ in range(args.pool - 5)] + + best, best_energy, trace = None, np.inf, [] + for round_index in range(args.rounds): + votes = np.zeros((size, size)) + for start in pool: + final = fast_pair_descent( + text_gpu, visual_gpu, + torch.from_numpy(np.ascontiguousarray(start)).to(device), + args.iterations, + ).cpu().numpy() + votes[np.arange(size), final] += 1.0 + + _, columns = linear_sum_assignment(-votes) + descended = fast_pair_descent( + text_gpu, visual_gpu, + torch.from_numpy(np.ascontiguousarray(columns)).to(device), + args.iterations, + ) + value = float(energy.energy(descended[None])[0]) + current = descended.cpu().numpy() + if value < best_energy: + best, best_energy = current, value + trace.append(value) + + # rebuild the pool around the current answer, keeping some of it random + # so a plateaued ladder still has a way out + fresh = int(args.fresh * args.pool) + pool = [current.copy()] + for _ in range(args.pool - fresh - 1): + perturbed = current.copy() + for _ in range(int(generator.integers(10, 120))): + i, j = generator.integers(0, size, 2) + perturbed[[i, j]] = perturbed[[j, i]] + pool.append(perturbed) + pool += [generator.permutation(size) for _ in range(fresh)] + return best, best_energy, trace + + +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) + visual_gpu = standardized(torch.from_numpy(visual).to(device)).double() + print(f"\n=== {label} (N={size}, chance={1/size:.4f})", flush=True) + + chosen, oracle_best = [], [] + for trial in range(args.trials): + hidden = np.random.default_rng(trial).permutation(size) + shuffled = text[np.ix_(hidden, hidden)] + text_gpu = standardized(torch.from_numpy(shuffled).to(device)).double() + energy = ClosedFormEnergy(text_gpu, visual_gpu, 1.0, 0.0, 256) + + results = [] + for index in range(args.ladders): + generator = np.random.default_rng(1000 * trial + index) + permutation, value, _ = ladder( + text_gpu, visual_gpu, visual, shuffled, energy, + args, generator, device, + ) + accuracy = float((hidden[permutation] == np.arange(size)).mean()) + results.append((value, accuracy)) + # selection by energy only -- the hidden pairing is never consulted + picked = min(results, key=lambda item: item[0])[1] + chosen.append(picked) + oracle_best.append(max(r[1] for r in results)) + print(f" trial {trial}: ladders={[round(a, 3) for _, a in results]} " + f"picked-by-energy={picked:.4f}", flush=True) + + row = { + "field": label, "size": size, "chance": 1.0 / size, + "accuracy_selected_by_energy": float(np.mean(chosen)), + "accuracy_if_oracle_picked_best_ladder": float(np.mean(oracle_best)), + } + rows.append(row) + print(f" MEAN selected-by-energy={row['accuracy_selected_by_energy']:.4f} " + f"(oracle ladder choice would give " + f"{row['accuracy_if_oracle_picked_best_ladder']:.4f})", flush=True) + + write_json(args.output, { + "protocol": ( + "Blind. Hidden permutation generated per trial and read only for " + "scoring; the ladder is chosen by energy, never by accuracy." + ), + "pool": args.pool, "rounds": args.rounds, "ladders": args.ladders, + "rows": rows, + }) + print(json.dumps({"done": True})) + + +if __name__ == "__main__": + main() 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() diff --git a/worldalign/match_battery.py b/worldalign/match_battery.py index 45fa26e..9035e62 100644 --- a/worldalign/match_battery.py +++ b/worldalign/match_battery.py @@ -123,12 +123,213 @@ def solve_bapg(visual, shuffled, **_): return coupling_to_permutation(coupling) +def solve_faq(visual, shuffled, **_): + """Frank-Wolfe on the Birkhoff polytope: the canonical QAP relaxation. + + This is the convex-concave family the register listed and never reached. + It optimises over doubly stochastic matrices and projects at the end, + which is a different relaxation from both the spectral solvers and GW. + """ + from scipy.optimize import quadratic_assignment + + result = quadratic_assignment( + visual, shuffled, method="faq", options={"maximize": True, "n_init": 1} + ) + return result.col_ind + + +def solve_faq_multi(visual, shuffled, **_): + """FAQ from many random doubly stochastic starts, best by its own objective.""" + from scipy.optimize import quadratic_assignment + + result = quadratic_assignment( + visual, shuffled, method="faq", + options={"maximize": True, "n_init": 30, "rng": 0}, + ) + return result.col_ind + + +def solve_2opt(visual, shuffled, **_): + from scipy.optimize import quadratic_assignment + + result = quadratic_assignment( + visual, shuffled, method="2opt", options={"maximize": True, "rng": 0} + ) + return result.col_ind + + +def solve_gw_multi(visual, shuffled, **_): + """GW is non-convex; restart it and keep the best by the GW objective.""" + import ot + + size = len(visual) + weights = ot.unif(size) + best, best_value = None, np.inf + generator = np.random.default_rng(0) + for attempt in range(12): + init = None + if attempt: + noise = generator.random((size, size)) + 1e-3 + init = noise / noise.sum() + try: + coupling, log = ot.gromov.gromov_wasserstein( + visual, shuffled, weights, weights, "square_loss", + G0=init, max_iter=200, tol_rel=1e-9, log=True, + ) + except Exception: + continue + value = float(log["gw_dist"]) + if value < best_value: + best, best_value = coupling, value + if best is None: + raise RuntimeError("all GW restarts failed") + return coupling_to_permutation(best) + + +def solve_gw_kl(visual, shuffled, **_): + import ot + + size = len(visual) + weights = ot.unif(size) + shift = min(visual.min(), shuffled.min()) + coupling = ot.gromov.gromov_wasserstein( + visual - shift + 1e-3, shuffled - shift + 1e-3, weights, weights, + "kl_loss", max_iter=200, tol_rel=1e-9, + ) + return coupling_to_permutation(coupling) + + +def solve_gw_deep_anneal(visual, shuffled, **_): + """The winning solver, annealed further and started from the smooth end.""" + import ot + + size = len(visual) + weights = ot.unif(size) + coupling = None + schedule = (2e-1, 1e-1, 5e-2, 3e-2, 2e-2, 1e-2, 7e-3, 5e-3, 3e-3, 2e-3, 1e-3) + for epsilon in schedule: + try: + candidate = ot.gromov.entropic_gromov_wasserstein( + visual, shuffled, weights, weights, "square_loss", + epsilon=epsilon, G0=coupling, max_iter=300, tol=1e-9, solver="PGD", + ) + except Exception: + break + if not np.isfinite(candidate).all(): + break + coupling = candidate + if coupling is None: + raise RuntimeError("deep anneal diverged immediately") + return coupling_to_permutation(coupling) + + +def solve_semirelaxed(visual, shuffled, **_): + import ot + + size = len(visual) + coupling = ot.gromov.semirelaxed_gromov_wasserstein( + visual, shuffled, ot.unif(size), "square_loss", max_iter=200, + ) + return coupling_to_permutation(coupling) + + +def _birkhoff_frank_wolfe(visual, shuffled, mu, coupling, steps): + """Frank-Wolfe on the Birkhoff polytope for one concavity setting.""" + from scipy.optimize import linear_sum_assignment + + size = len(visual) + for step in range(steps): + gradient = 2.0 * (visual @ coupling @ shuffled) + 2.0 * mu * coupling + _, columns = linear_sum_assignment(-gradient) + vertex = np.zeros_like(coupling) + vertex[np.arange(size), columns] = 1.0 + direction = vertex - coupling + # exact line search on the quadratic along the segment + quad = float((direction * (visual @ direction @ shuffled)).sum()) + mu * float( + (direction * direction).sum() + ) + linear = float((direction * gradient).sum()) + if abs(quad) < 1e-12: + gamma = 1.0 if linear > 0 else 0.0 + else: + gamma = float(np.clip(-linear / (2.0 * quad), 0.0, 1.0)) if quad < 0 else ( + 1.0 if linear > 0 else 0.0 + ) + if gamma <= 1e-12: + break + coupling = coupling + gamma * direction + return coupling + + +def solve_path(visual, shuffled, **_): + """Convex-to-concave path following over the Birkhoff polytope. + + The relaxation the register named and never ran. Maximising a concave + surrogate over the polytope has one smooth interior optimum; maximising a + convex one attains its maximum at a vertex, which is a permutation. Adding + mu*||P||^2 and sweeping mu from negative to positive walks continuously + between the two, and the argument for it is exactly our situation: the + smooth end has no local optima to be trapped by, and the solution is + tracked into the sharp end rather than searched for there. + """ + size = len(visual) + scale = float(np.abs(np.linalg.eigvalsh(visual)).max() + * np.abs(np.linalg.eigvalsh(shuffled)).max()) + coupling = np.full((size, size), 1.0 / size) + for fraction in np.linspace(-1.0, 1.0, 21): + coupling = _birkhoff_frank_wolfe( + visual, shuffled, fraction * scale, coupling, 30 + ) + return coupling_to_permutation(coupling) + + +def solve_moment_ladder(visual, shuffled, **_): + """Match the eigenvalue moments of M, not just its sum. + + The energy is a moment functional in disguise: the pairwise term is the + first moment of M = permuted-text * visual and the triangle term its + third. The sequence {tr(M^k)} is what pins down M's spectrum, and we have + only ever used two of its entries. This optimises a weighted combination + over several k by Frank-Wolfe, which is the cheapest way to ask whether a + longer moment sequence has a wider basin than a single term. + """ + from scipy.optimize import linear_sum_assignment + + size = len(visual) + coupling = np.full((size, size), 1.0 / size) + for step in range(120): + permuted = coupling @ shuffled @ coupling.T + product = permuted * visual + # d/dP of tr(M^k) terms, accumulated over the ladder + gradient = 2.0 * (visual @ coupling @ shuffled) + squared = product @ product + gradient = gradient + 1.0 * ( + (visual * product) @ coupling @ shuffled + + (visual * squared) @ coupling @ shuffled + ) + _, columns = linear_sum_assignment(-gradient) + vertex = np.zeros_like(coupling) + vertex[np.arange(size), columns] = 1.0 + gamma = 2.0 / (step + 2.0) + coupling = coupling + gamma * (vertex - coupling) + return coupling_to_permutation(coupling) + + SOLVERS = { "grampa (current)": solve_grampa, "umeyama": solve_umeyama, "gromov-wasserstein": solve_gw, "entropic GW annealed": solve_gw_annealed, "BAPG GW": solve_bapg, + "FAQ (Frank-Wolfe)": solve_faq, + "FAQ x30 restarts": solve_faq_multi, + "2-opt": solve_2opt, + "GW x12 restarts": solve_gw_multi, + "GW kl-loss": solve_gw_kl, + "entropic GW deep anneal": solve_gw_deep_anneal, + "semirelaxed GW": solve_semirelaxed, + "PATH convex-concave": solve_path, + "moment ladder tr(M^k)": solve_moment_ladder, } |
