summaryrefslogtreecommitdiff
path: root/worldalign/synth_deep_gate.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/synth_deep_gate.py')
-rw-r--r--worldalign/synth_deep_gate.py203
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()