diff options
Diffstat (limited to 'worldalign/synth_recovery.py')
| -rw-r--r-- | worldalign/synth_recovery.py | 184 |
1 files changed, 184 insertions, 0 deletions
diff --git a/worldalign/synth_recovery.py b/worldalign/synth_recovery.py new file mode 100644 index 0000000..e045a72 --- /dev/null +++ b/worldalign/synth_recovery.py @@ -0,0 +1,184 @@ +"""Blind recovery on synthetic set-kernel fields: the end-to-end test. + +Builds the connected-component descriptor field and the phrase +bag-of-words field, hides the text order behind a shuffle, and runs +parallel tempering on the relational energy alone. Recovery accuracy +against the hidden truth is the first end-to-end measurement of world +matching in the closed world. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +from tqdm import tqdm + +from .blind_recovery import arm_tempering, permutation_energy_batch +from .common import read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .synth_cc_battery import ( + component_descriptors, + moment_field, + onehot_descriptors, + phrase_bow_sets, +) +from .synth_set_battery import set_similarity_field +from .synth_towers import load_image + + +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("--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("--unary-weight", type=float, default=0.0) + parser.add_argument("--init", default="random") + parser.add_argument("--residualize-size", action="store_true", default=False) + parser.add_argument( + "--features", choices=["descriptors", "onehot"], default="onehot" + ) + parser.add_argument("--kernel", choices=["matching", "moment"], default="moment") + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/recovery_end_to_end.json" + ) + return parser.parse_args() + + +def residualize(field: torch.Tensor, sizes: torch.Tensor) -> torch.Tensor: + """Regress the set-size nuisance out of a matching-value field. + + Set sizes are unimodal observables; their sum, difference, and product + explain a size-driven component that differs between modalities and + is exploitable by counterfeit assignments. + """ + n = len(field) + mask = ~torch.eye(n, dtype=torch.bool) + features = torch.stack( + [ + (sizes[:, None] + sizes[None, :])[mask], + (sizes[:, None] - sizes[None, :]).abs()[mask], + (sizes[:, None] * sizes[None, :])[mask], + torch.ones(int(mask.sum()), dtype=torch.float64), + ], + dim=1, + ) + values = field.double()[mask] + solution = torch.linalg.lstsq(features, values[:, None]).solution + residual = values - (features @ solution).squeeze(1) + output = field.double().clone() + output[mask] = residual + output.fill_diagonal_(0.0) + return output.float() + + +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"]) + + import torch.nn.functional as F + + per_view_fields = [] + for view in range(args.vision_views): + raw_sets = [] + for row in tqdm(rows, desc=f"cc v{view}"): + image = load_image(image_dir / f"scene{row:06d}_v{view}.png") + sprites, _ = component_descriptors(image, args.merge_distance) + raw_sets.append(sprites) + if args.features == "onehot": + vision_sets = [ + F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets) + ] + else: + vision_sets = [F.normalize(v, dim=-1) for v in raw_sets] + if args.kernel == "moment": + per_view_fields.append(moment_field(vision_sets)) + else: + per_view_fields.append(set_similarity_field(vision_sets)) + visual_field = torch.stack(per_view_fields).mean(0) + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + if args.kernel == "moment": + text_field = moment_field(text_sets) + else: + text_field = set_similarity_field(text_sets) + + if args.residualize_size: + vision_sizes = torch.tensor( + [len(s) for s in vision_sets], dtype=torch.float64 + ) + text_sizes = torch.tensor([len(s) for s in text_sets], dtype=torch.float64) + visual_field = residualize(visual_field, vision_sizes) + text_field = residualize(text_field, text_sizes) + + 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 = text_field[hidden][:, hidden] + + text_standardized = standardize_relation(text_input.double())[0].float().to(device) + visual_standardized = ( + standardize_relation(visual_field.double())[0].float().to(device) + ) + true_energy = float( + permutation_energy_batch( + text_standardized, visual_standardized, truth[None].to(device) + )[0] + ) + report = { + "protocol": ( + "Set-kernel fields from released renders and captions; text " + "order hidden behind a shuffle; tempering sees no truth. " + "Hidden truth scores the outcome only." + ), + "split": args.split, + "samples": size, + "true_energy": true_energy, + "chance_accuracy": 1.0 / size, + "tempering": arm_tempering( + text_standardized, + visual_standardized, + truth, + true_energy, + args, + generator, + unary=None, + ), + } + best = max( + report["tempering"]["replicas"] + [report["tempering"]["best"]], + key=lambda c: c["accuracy"], + ) + by_energy = min( + report["tempering"]["replicas"] + [report["tempering"]["best"]], + key=lambda c: c["energy"], + ) + report["summary"] = { + "best_accuracy": best["accuracy"], + "best_accuracy_energy_over_true": best["energy_over_true"], + "lowest_energy_accuracy": by_energy["accuracy"], + "lowest_energy_over_true": by_energy["energy_over_true"], + } + print(json.dumps({"summary": report["summary"]})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
