diff options
| author | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
|---|---|---|
| committer | Yuren Hao <blackhao0426@gmail.com> | 2026-08-01 14:10:03 -0500 |
| commit | a62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch) | |
| tree | ee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign/natural_fields.py | |
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built
within each modality and are invariant to how each side labels its own
features; the cross-modal bridge is a coupling searched under an energy
that is a closed-form functional of one matrix; solving is spectral
initialisation followed by exact local refinement.
Evidence: in a procedurally generated closed world, blind recovery of a
hidden image-caption correspondence reaches 95.3% at 256 scenes against
0.39% chance, and the recovered pairs transfer to 200 held-out scenes at
93.0% exact retrieval with random-pair and shuffled-image controls at or
near chance. Cross-modal value correspondence is derived from disjoint
corpora rather than declared. On Visual Genome the field correlation
reaches 0.656 against the 0.9 that polynomial recovery needs, with the
deficit attributed away from segmentation and discretisation.
Protocol: no image-text pair enters any objective, optimiser,
initialisation, or model selection; hidden pairs score orderings only.
Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/natural_fields.py')
| -rw-r--r-- | worldalign/natural_fields.py | 224 |
1 files changed, 224 insertions, 0 deletions
diff --git a/worldalign/natural_fields.py b/worldalign/natural_fields.py new file mode 100644 index 0000000..7221a3a --- /dev/null +++ b/worldalign/natural_fields.py @@ -0,0 +1,224 @@ +"""Natural-data relation fields under the Tier 0 recipe. + +Each side is encoded into its own discovered factor coordinates -- text +word classes from distributional induction, vision segment classes from +feature clustering -- and the two coordinate systems are paired by joint +structure, never by a declared lexicon. Scene states are the resulting +sets of segment or phrase codes; relations are moment-kernel similarities +between scenes. + +The field correlation at the hidden pairing is the go/no-go statistic: +polynomial recovery needs roughly 0.9, and the synthetic world showed +nothing works below it. Hidden pairs are read only to compute it. +""" + +from __future__ import annotations + +import argparse +import json +import re +from collections import Counter, defaultdict +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy.optimize import linear_sum_assignment +from sklearn.cluster import KMeans + +from .common import read_json, seed_everything, write_json +from .natural_families import distributional_families, load_phrases, statistics +from .synth_cc_battery import moment_field + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--vg-dir", default="artifacts/vg_5k") + parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt") + parser.add_argument("--samples", type=int, default=256) + parser.add_argument("--fit-nodes", type=int, default=2000) + parser.add_argument("--vision-classes", type=int, default=24) + parser.add_argument("--text-classes", type=int, default=24) + parser.add_argument("--clusters", type=int, default=24) + parser.add_argument("--min-count", type=int, default=200) + parser.add_argument("--max-vocabulary", type=int, default=300) + parser.add_argument("--context-window", type=int, default=2) + parser.add_argument("--min-family", type=int, default=3) + parser.add_argument("--components", type=int, default=64) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--output", default="artifacts/vg_5k/natural_fields.pt") + return parser.parse_args() + + +def text_scene_codes( + records: list[dict], word_class: dict[str, int], classes: int +) -> dict[str, torch.Tensor]: + """One code vector per region description: its word-class profile.""" + codes: dict[str, list[torch.Tensor]] = {} + for record in records: + vectors = [] + for phrase in record["region_closed"]: + vector = torch.zeros(classes) + for token in re.findall(r"[a-z]+", phrase.lower()): + if token in word_class: + vector[word_class[token]] += 1.0 + if vector.sum() > 0: + vectors.append(vector) + if vectors: + codes[record["node_id"]] = F.normalize(torch.stack(vectors), dim=-1) + return codes + + +def vision_scene_codes( + state: dict, model: KMeans, classes: int +) -> dict[str, torch.Tensor]: + """One code vector per segment: its feature-class assignment.""" + codes: dict[str, torch.Tensor] = {} + for node, segments in zip(state["node_ids"], state["segments"]): + if not segments: + continue + labels = model.predict( + np.stack([segment["feature"] for segment in segments]).astype(np.float64) + ) + vectors = torch.zeros(len(segments), classes) + for row, label in enumerate(labels): + vectors[row, int(label)] = 1.0 + codes[node] = F.normalize(vectors, dim=-1) + return codes + + +def align_classes( + text_codes: dict[str, torch.Tensor], + vision_codes: dict[str, torch.Tensor], + classes: int, +) -> np.ndarray: + """Pair vision classes to text classes by marginal frequency rank. + + Within-scene co-occurrence is the stronger signal but needs a second + factor to condition on; frequency is the available unimodal statistic + at this stage and is reported as the first pass. + """ + text_mass = torch.zeros(classes) + for code in text_codes.values(): + text_mass += code.sum(0) + vision_mass = torch.zeros(classes) + for code in vision_codes.values(): + vision_mass += code.sum(0) + text_order = torch.argsort(text_mass, descending=True).numpy() + vision_order = torch.argsort(vision_mass, descending=True).numpy() + mapping = np.empty(classes, dtype=int) + mapping[vision_order] = text_order + return mapping + + +def moment_states(codes: torch.Tensor) -> torch.Tensor: + first = codes.mean(0) + second = (codes[:, :, None] * codes[:, None, :]).mean(0).flatten() + return torch.cat([first, second]) + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + vg_dir = Path(args.vg_dir) + text_records = [ + json.loads(line) + for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines() + if line.strip() + ] + truth = [ + json.loads(line) + for line in (vg_dir / "ground_truth.private.jsonl") + .read_text(encoding="utf-8") + .splitlines() + if line.strip() + ] + state = torch.load(args.objects, map_location="cpu", weights_only=False) + + phrases = load_phrases(vg_dir / "text_nodes.jsonl", "region_closed") + counts = Counter(token for tokens in phrases for token in set(tokens)) + vocabulary = [ + word + for word, count in counts.most_common(args.max_vocabulary) + if count >= args.min_count + ] + _, _, context = statistics(phrases, set(vocabulary), args.context_window) + args.clusters = args.text_classes + families = distributional_families(vocabulary, context, args) + word_class = { + word: index for index, family in enumerate(families) for word in family + } + text_classes = len(families) + + features = np.stack( + [ + segment["feature"] + for segments in state["segments"][: args.fit_nodes] + for segment in segments + ] + ).astype(np.float64) + vision_model = KMeans( + args.vision_classes, n_init=10, random_state=args.seed + ).fit(features) + + text_codes = text_scene_codes(text_records, word_class, text_classes) + vision_codes = vision_scene_codes(state, vision_model, args.vision_classes) + + shared = min(text_classes, args.vision_classes) + mapping = align_classes( + {k: v[:, :shared] for k, v in text_codes.items()}, + {k: v[:, :shared] for k, v in vision_codes.items()}, + shared, + ) + + pairs = [ + pair + for pair in truth + if pair["vision_node_id"] in vision_codes and pair["text_node_id"] in text_codes + ][: args.samples] + vision_sets, text_sets = [], [] + for pair in pairs: + vision = vision_codes[pair["vision_node_id"]][:, :shared] + remapped = torch.zeros_like(vision) + for source in range(shared): + remapped[:, mapping[source]] = vision[:, source] + vision_sets.append(F.normalize(remapped, dim=-1)) + text_sets.append( + F.normalize(text_codes[pair["text_node_id"]][:, :shared], dim=-1) + ) + + visual_field = moment_field(vision_sets) + text_field = moment_field(text_sets) + size = len(pairs) + mask = ~np.eye(size, dtype=bool) + correlation = float( + np.corrcoef( + visual_field.double().numpy()[mask], text_field.double().numpy()[mask] + )[0, 1] + ) + torch.save( + { + "visual_field": visual_field, + "text_field": text_field, + "pairs": [p["vision_node_id"] for p in pairs], + }, + args.output, + ) + summary = { + "protocol": ( + "Text classes from distributional induction, vision classes " + "from segment-feature clustering, paired by marginal frequency " + "rank. Hidden pairs are read only for the correlation." + ), + "samples": size, + "text_classes": text_classes, + "vision_classes": args.vision_classes, + "field_correlation_at_truth": correlation, + "go_no_go": "recovery needs about 0.9", + } + print(json.dumps(summary)) + write_json(str(args.output).replace(".pt", ".json"), summary) + + +if __name__ == "__main__": + main() |
