diff options
Diffstat (limited to 'worldalign/synth_deep_gate.py')
| -rw-r--r-- | worldalign/synth_deep_gate.py | 203 |
1 files changed, 203 insertions, 0 deletions
diff --git a/worldalign/synth_deep_gate.py b/worldalign/synth_deep_gate.py new file mode 100644 index 0000000..74bfea9 --- /dev/null +++ b/worldalign/synth_deep_gate.py @@ -0,0 +1,203 @@ +"""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() |
