diff options
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() |
