summaryrefslogtreecommitdiff
path: root/worldalign/synth_precision_gate.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/synth_precision_gate.py')
-rw-r--r--worldalign/synth_precision_gate.py279
1 files changed, 279 insertions, 0 deletions
diff --git a/worldalign/synth_precision_gate.py b/worldalign/synth_precision_gate.py
new file mode 100644
index 0000000..46887fa
--- /dev/null
+++ b/worldalign/synth_precision_gate.py
@@ -0,0 +1,279 @@
+"""Channel-matched (precision-weighted) relational energy for the synth world.
+
+Planted-problem theory: search is glassy when the energy mismatches the
+generative channel. The orbit provides the channel unimodally -- the
+variance of each vision relation entry across re-rendered views measures
+exactly how corrupted that entry is (segmentation errors are the dominant
+noise and are view-dependent), so entries are weighted by their orbit
+precision. Text fields are parse-deterministic and enter unweighted.
+
+The weighted energy loses the all-pairs closed form (the quadratic terms
+no longer cancel), but each swap delta is still O(N) and vectorizes over
+sampled pairs; gates and tempering below use that form. Hidden pairs
+score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_cc_battery import (
+ component_descriptors,
+ moment_field,
+ onehot_descriptors,
+ phrase_bow_sets,
+)
+
+
+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("--precision-floor", type=float, default=1e-4)
+ parser.add_argument("--random-perms", type=int, default=300)
+ parser.add_argument("--transposition-samples", type=int, default=100000)
+ parser.add_argument("--descent-restarts", type=int, default=5)
+ parser.add_argument("--descent-max-steps", type=int, default=5000)
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--tempering-rounds", type=int, default=60000)
+ parser.add_argument("--temp-high", type=float, default=3e-3)
+ parser.add_argument("--temp-low", type=float, default=1e-5)
+ 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/precision_gate.json"
+ )
+ return parser.parse_args()
+
+
+def standardized(field: torch.Tensor) -> torch.Tensor:
+ mask = ~torch.eye(len(field), dtype=torch.bool)
+ values = field[mask]
+ out = (field - values.mean()) / values.std().clamp_min(1e-9)
+ return out.masked_fill(~mask, 0.0)
+
+
+def weighted_energy(
+ text: torch.Tensor, visual: torch.Tensor, weights: torch.Tensor,
+ permutations: torch.Tensor,
+) -> torch.Tensor:
+ fields = text[permutations[:, :, None], permutations[:, None, :]]
+ difference = (fields - visual) ** 2 * weights
+ mask = ~torch.eye(text.shape[-1], dtype=torch.bool, device=text.device)
+ return difference[:, mask].sum(-1) / weights[mask].sum()
+
+
+def swap_deltas(
+ permuted_text: torch.Tensor,
+ visual: torch.Tensor,
+ weights: torch.Tensor,
+ pairs_p: torch.Tensor,
+ pairs_q: torch.Tensor,
+) -> torch.Tensor:
+ """Weighted swap deltas, O(N) per proposal, batched over replicas.
+
+ permuted_text: [R, N, N]; pairs: [R, P]. Row and column contributions
+ are equal by symmetry of all three matrices; the k in {p, q} terms are
+ excluded because the (p, q) entry itself is unchanged by the swap.
+ """
+ replica_index = torch.arange(len(permuted_text), device=permuted_text.device)
+ rows_p = permuted_text[replica_index[:, None], pairs_p]
+ rows_q = permuted_text[replica_index[:, None], pairs_q]
+ visual_p, visual_q = visual[pairs_p], visual[pairs_q]
+ weight_p, weight_q = weights[pairs_p], weights[pairs_q]
+ new_p = (rows_q - visual_p) ** 2 * weight_p
+ old_p = (rows_p - visual_p) ** 2 * weight_p
+ new_q = (rows_p - visual_q) ** 2 * weight_q
+ old_q = (rows_q - visual_q) ** 2 * weight_q
+ total = (new_p - old_p + new_q - old_q).sum(-1)
+ columns = torch.stack([pairs_p, pairs_q], -1)
+ correction = torch.zeros_like(total)
+ for slot in range(2):
+ chosen = columns[..., slot]
+ correction = correction + (
+ (rows_q.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2
+ - (rows_p.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2
+ ).squeeze(-1) * weight_p.gather(-1, chosen[..., None]).squeeze(-1)
+ correction = correction + (
+ (rows_p.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2
+ - (rows_q.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2
+ ).squeeze(-1) * weight_q.gather(-1, chosen[..., None]).squeeze(-1)
+ mask_count = weights[~torch.eye(len(visual), dtype=torch.bool, device=visual.device)].sum()
+ return 2.0 * (total - correction) / mask_count
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ captions = read_json(Path(args.data_dir, "captions.json"))["captions"]
+ rows = manifest[args.split][: args.samples]
+ image_dir = Path(manifest["image_dir"])
+ from .synth_towers import load_image
+
+ per_view = []
+ for view in range(args.vision_views):
+ raw = []
+ for row in tqdm(rows, desc=f"cc v{view}"):
+ sprites, _ = component_descriptors(
+ load_image(image_dir / f"scene{row:06d}_v{view}.png"),
+ args.merge_distance,
+ )
+ raw.append(sprites)
+ sets = [F.normalize(v, dim=-1) for v in onehot_descriptors(raw)]
+ per_view.append(moment_field(sets))
+ stack = torch.stack(per_view)
+ visual_field = stack.mean(0)
+ variance = stack.var(0)
+ weights = 1.0 / (variance + args.precision_floor)
+ weights = weights / weights.mean()
+
+ text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"])
+ text_field = moment_field(text_sets)
+
+ 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)
+ text_input = standardized(text_field[hidden][:, hidden].double()).float().to(device)
+ visual = standardized(visual_field.double()).float().to(device)
+ weight_map = weights.float().to(device)
+
+ identity = torch.arange(size, device=device)
+ true_energy = float(
+ weighted_energy(text_input, visual, weight_map, truth[None].to(device))[0]
+ )
+ random_perms = torch.stack(
+ [torch.argsort(torch.rand(size, generator=generator)) for _ in range(args.random_perms)]
+ ).to(device)
+ random_energies = weighted_energy(text_input, visual, weight_map, random_perms)
+ gate_a = {
+ "true": true_energy,
+ "random_mean": float(random_energies.mean()),
+ "random_std": float(random_energies.std()),
+ "true_z": float((random_energies.mean() - true_energy) / random_energies.std().clamp_min(1e-12)),
+ }
+
+ samples = min(args.transposition_samples, size * (size - 1) // 2)
+ pairs_p = torch.randint(0, size, (1, samples), generator=generator).to(device)
+ pairs_q = torch.randint(0, size, (1, samples), generator=generator).to(device)
+ valid = (pairs_p != pairs_q).squeeze(0)
+ fields_true = text_input[truth.to(device)][:, truth.to(device)][None]
+ deltas = swap_deltas(fields_true, visual, weight_map, pairs_p, pairs_q).squeeze(0)[valid]
+ gate_b = {
+ "sampled": int(valid.sum()),
+ "improving_fraction": float((deltas < 0).float().mean()),
+ }
+
+ def descent(start: torch.Tensor) -> dict:
+ current = start.clone()
+ energy = float(weighted_energy(text_input, visual, weight_map, current[None])[0])
+ for _ in range(args.descent_max_steps):
+ fields = text_input[current][:, current][None]
+ cp = torch.randint(0, size, (1, 4096), generator=generator).to(device)
+ cq = torch.randint(0, size, (1, 4096), generator=generator).to(device)
+ dd = swap_deltas(fields, visual, weight_map, cp, cq).squeeze(0)
+ best = int(dd.argmin())
+ if float(dd[best]) >= -1e-12:
+ break
+ p, q = int(cp[0, best]), int(cq[0, best])
+ current[[p, q]] = current[[q, p]]
+ energy += float(dd[best])
+ return {
+ "accuracy": float((current.cpu() == truth).float().mean()),
+ "energy": float(weighted_energy(text_input, visual, weight_map, current[None])[0]),
+ }
+
+ from_true = descent(truth.to(device))
+ restarts = [
+ descent(torch.argsort(torch.rand(size, generator=generator)).to(device))
+ for _ in range(args.descent_restarts)
+ ]
+ best_random = min(r["energy"] for r in restarts)
+
+ # Tempering recovery with the weighted deltas.
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ args.replicas,
+ ).to(device)
+ permutations = torch.stack(
+ [torch.randperm(size, generator=generator).to(device) for _ in range(args.replicas)]
+ )
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ for round_index in range(args.tempering_rounds):
+ fields = text_input[permutations[:, :, None], permutations[:, None, :]]
+ cp = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device)
+ cq = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device)
+ dd = swap_deltas(fields, visual, weight_map, cp, cq)
+ noise = torch.rand(args.replicas, 24, generator=generator).to(device)
+ ok = (dd < -temperatures[:, None] * noise.clamp_min(1e-12).log()) & (cp != cq)
+ for replica in range(args.replicas):
+ hits = torch.nonzero(ok[replica])
+ if not len(hits):
+ continue
+ first = int(hits[0, 0])
+ p, q = int(cp[replica, first]), int(cq[replica, first])
+ permutations[replica][[p, q]] = permutations[replica][[q, p]]
+ energies[replica] = energies[replica] + dd[replica, first]
+ if round_index % args.exchange_every == 0:
+ for replica in range(args.replicas - 1):
+ gap = (energies[replica] - energies[replica + 1]) * (
+ 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1]
+ )
+ if gap > 0 or torch.rand(1, generator=generator).item() < float(gap.exp()):
+ permutations[[replica, replica + 1]] = permutations[[replica + 1, replica]]
+ energies[[replica, replica + 1]] = energies[[replica + 1, replica]]
+ if round_index % 10000 == 0:
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ accuracies = (permutations.cpu() == truth[None]).float().mean(-1)
+ cold = int(energies.argmin())
+
+ report = {
+ "protocol": (
+ "Relation entries are weighted by orbit-derived precision "
+ "(variance of the vision field across re-rendered views); "
+ "weights are unimodal statistics. Hidden pairs score only."
+ ),
+ "samples": size,
+ "weight_stats": {
+ "min": float(weights.min()),
+ "median": float(weights.median()),
+ "max": float(weights.max()),
+ },
+ "gate_a": gate_a,
+ "gate_b": gate_b,
+ "descent_from_true": from_true,
+ "descent_from_random": restarts,
+ "tempering": {
+ "cold_energy": float(energies[cold]),
+ "cold_accuracy": float(accuracies[cold]),
+ "best_accuracy": float(accuracies.max()),
+ },
+ "verdict": {
+ "true_energy": true_energy,
+ "best_random_descent": best_random,
+ "counterfeit_found": bool(best_random < true_energy and min(r["accuracy"] for r in restarts) < 0.5),
+ "recovery_accuracy": float(accuracies.max()),
+ },
+ }
+ print(json.dumps({"verdict": report["verdict"], "gate_b": gate_b, "from_true": from_true}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()