diff options
Diffstat (limited to 'worldalign/synth_triangle_recovery.py')
| -rw-r--r-- | worldalign/synth_triangle_recovery.py | 173 |
1 files changed, 173 insertions, 0 deletions
diff --git a/worldalign/synth_triangle_recovery.py b/worldalign/synth_triangle_recovery.py new file mode 100644 index 0000000..a9bca61 --- /dev/null +++ b/worldalign/synth_triangle_recovery.py @@ -0,0 +1,173 @@ +"""Blind recovery under the triangle energy. + +The triangle gate holds the true assignment at N=512 with no counterfeit +below it. The search question is separate and now worth asking again: +tempering with third-order swap deltas, from random starts, with the +hidden order behind a shuffle. Recovery accuracy is the end-to-end world +matching measurement. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch + +from .common import read_json, seed_everything, write_json +from .synth_triangle_gate import TriangleEnergy, build_fields, standardized + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--triples", type=int, default=400000) + parser.add_argument("--pair-weight", type=float, default=1.0) + parser.add_argument("--triangle-weight", type=float, default=1.0) + parser.add_argument("--replicas", type=int, default=6) + parser.add_argument("--rounds", type=int, default=4000) + parser.add_argument("--proposals", type=int, default=48) + parser.add_argument("--temp-high", type=float, default=2e-2) + parser.add_argument("--temp-low", type=float, default=1e-4) + parser.add_argument("--exchange-every", type=int, default=20) + parser.add_argument("--greedy-rounds", type=int, default=3000) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/triangle_recovery.json" + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + manifest = read_json(Path(args.data_dir, "manifest.json")) + rows = manifest[args.split][: args.samples] + visual_field, text_field = build_fields(args, rows, manifest) + + device = torch.device(args.device) + size = len(rows) + generator = torch.Generator().manual_seed(args.seed) + hidden = torch.randperm(size, generator=generator) + truth = torch.argsort(hidden).to(device) + text = standardized(text_field[hidden][:, hidden].double().to(device)).float() + visual = standardized(visual_field.double().to(device)).float() + + triples = torch.randint(0, size, (args.triples, 3), generator=generator) + triples = triples[ + (triples[:, 0] != triples[:, 1]) + & (triples[:, 1] != triples[:, 2]) + & (triples[:, 0] != triples[:, 2]) + ].to(device) + energy = TriangleEnergy( + text, visual, triples, args.pair_weight, args.triangle_weight + ) + true_energy = energy.total(truth) + + temperatures = torch.logspace( + torch.log10(torch.tensor(args.temp_low)), + torch.log10(torch.tensor(args.temp_high)), + args.replicas, + ) + states = [ + torch.argsort(torch.rand(size, generator=generator)).to(device) + for _ in range(args.replicas) + ] + energies = [energy.total(state) for state in states] + history = [] + for round_index in range(args.rounds): + for replica in range(args.replicas): + temperature = float(temperatures[replica]) + for _ in range(args.proposals): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(states[replica], p, q) + threshold = -temperature * float( + torch.rand(1, generator=generator).clamp_min(1e-12).log() + ) + if delta < threshold: + states[replica][[p, q]] = states[replica][[q, p]] + energies[replica] += delta + if round_index % args.exchange_every == 0: + for replica in range(args.replicas - 1): + gap = (energies[replica] - energies[replica + 1]) * ( + 1.0 / float(temperatures[replica]) + - 1.0 / float(temperatures[replica + 1]) + ) + if gap > 0 or float(torch.rand(1, generator=generator)) < min( + 1.0, float(torch.tensor(gap).exp()) + ): + states[replica], states[replica + 1] = ( + states[replica + 1], + states[replica], + ) + energies[replica], energies[replica + 1] = ( + energies[replica + 1], + energies[replica], + ) + if round_index % 200 == 0: + cold = min(range(args.replicas), key=lambda r: energies[r]) + record = { + "round": round_index, + "cold_energy": energies[cold], + "cold_accuracy": float( + (states[cold].cpu() == truth.cpu()).float().mean() + ), + "energy_over_true": energies[cold] / true_energy - 1.0, + } + history.append(record) + print(json.dumps(record)) + + # Greedy polish of the coldest replica. + cold = min(range(args.replicas), key=lambda r: energies[r]) + current = states[cold].clone() + for _ in range(args.greedy_rounds): + best_delta, best_pair = 0.0, None + for _ in range(96): + p = int(torch.randint(0, size, (1,), generator=generator)) + q = int(torch.randint(0, size, (1,), generator=generator)) + if p == q: + continue + delta = energy.swap_delta(current, p, q) + if delta < best_delta: + best_delta, best_pair = delta, (p, q) + if best_pair is None: + break + p, q = best_pair + current[[p, q]] = current[[q, p]] + + final = { + "accuracy": float((current.cpu() == truth.cpu()).float().mean()), + "energy": energy.total(current), + } + final["energy_over_true"] = final["energy"] / true_energy - 1.0 + report = { + "protocol": ( + "Tempering on the pairwise-plus-triangle energy from random " + "starts; text order hidden behind a shuffle; hidden truth " + "scores the outcome only." + ), + "samples": size, + "true_energy": true_energy, + "chance_accuracy": 1.0 / size, + "history": history, + "final_polished": final, + "replica_accuracies": [ + float((state.cpu() == truth.cpu()).float().mean()) for state in states + ], + } + print(json.dumps({"final": final})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
