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