"""Corrected gate: the deepest reachable minimum, not descent retention. Today's lesson: descent retention measures the search operator, not the energy. Sampled-proposal descent keeps the truth for every candidate energy, while long tempering on the same energy reaches states well below it. The only decision-relevant question is therefore E(truth) <= E(deepest state a strong searcher reaches) ? This module answers it uniformly for the candidate energies (pairwise moment kernel, triangle-only, and their sum) with one strong searcher: long parallel tempering with large proposal batches from random starts, plus a truth-initialized tempering arm that reports whether the truth itself survives thermal agitation. Hidden pairs score only. """ 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=256) 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=200000) parser.add_argument( "--energies", default="pair,triangle,both", help="Comma list from pair, triangle, both.", ) parser.add_argument("--replicas", type=int, default=6) parser.add_argument("--rounds", type=int, default=1500) parser.add_argument("--proposals", type=int, default=64) parser.add_argument("--temp-high", type=float, default=3e-2) parser.add_argument("--temp-low", type=float, default=1e-4) parser.add_argument("--exchange-every", type=int, default=20) parser.add_argument("--device", default="cuda:3") parser.add_argument("--seed", type=int, default=20260731) parser.add_argument("--output", default="artifacts/synth_v0/deep_gate.json") return parser.parse_args() def temper( energy: TriangleEnergy, starts: list[torch.Tensor], truth: torch.Tensor, args: argparse.Namespace, generator: torch.Generator, label: str, ) -> dict: size = energy.size temperatures = torch.logspace( torch.log10(torch.tensor(args.temp_low)), torch.log10(torch.tensor(args.temp_high)), len(starts), ) states = [start.clone() for start in starts] energies = [energy.total(state) for state in states] best = {"energy": min(energies), "accuracy": 0.0} for round_index in range(args.rounds): for replica in range(len(states)): 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(len(states) - 1): gap = (energies[replica] - energies[replica + 1]) * ( 1.0 / float(temperatures[replica]) - 1.0 / float(temperatures[replica + 1]) ) accept = gap > 0 or float( torch.rand(1, generator=generator) ) < min(1.0, float(torch.tensor(gap).exp())) if accept: states[replica], states[replica + 1] = ( states[replica + 1], states[replica], ) energies[replica], energies[replica + 1] = ( energies[replica + 1], energies[replica], ) cold = min(range(len(states)), key=lambda r: energies[r]) if energies[cold] < best["energy"]: best = { "energy": energies[cold], "accuracy": float( (states[cold].cpu() == truth.cpu()).float().mean() ), "round": round_index, } exact = [energy.total(state) for state in states] cold = min(range(len(states)), key=lambda r: exact[r]) return { "arm": label, "best_seen": best, "final_cold_energy": exact[cold], "final_cold_accuracy": float( (states[cold].cpu() == truth.cpu()).float().mean() ), "final_accuracies": [ float((state.cpu() == truth.cpu()).float().mean()) for state in states ], } 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) weights = { "pair": (1.0, 0.0), "triangle": (0.0, 1.0), "both": (1.0, 1.0), } report = { "protocol": ( "The decision statistic is the deepest energy a strong " "searcher reaches versus the energy of the truth. Descent " "retention is reported but not used: it measures the search " "operator. Hidden pairs score only." ), "samples": size, "triples": len(triples), "energies": {}, } for name in (item.strip() for item in args.energies.split(",")): pair_weight, triangle_weight = weights[name] energy = TriangleEnergy(text, visual, triples, pair_weight, triangle_weight) true_energy = energy.total(truth) random_starts = [ torch.argsort(torch.rand(size, generator=generator)).to(device) for _ in range(args.replicas) ] from_random = temper(energy, random_starts, truth, args, generator, "random") from_truth = temper( energy, [truth.clone() for _ in range(args.replicas)], truth, args, generator, "truth", ) deepest = min(from_random["best_seen"]["energy"], from_truth["best_seen"]["energy"]) entry = { "true_energy": true_energy, "from_random": from_random, "from_truth": from_truth, "deepest_seen": deepest, "margin_over_true": deepest / true_energy - 1.0, "passes": bool(deepest >= true_energy - 1e-9), "recovery_accuracy": max( from_random["best_seen"]["accuracy"], from_random["final_cold_accuracy"], ), } report["energies"][name] = entry print(json.dumps({name: {k: entry[k] for k in ("true_energy", "deepest_seen", "margin_over_true", "passes", "recovery_accuracy")}})) write_json(args.output, report) print(f"Wrote {args.output}") if __name__ == "__main__": main()