"""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()