summaryrefslogtreecommitdiff
path: root/worldalign/natural_fields.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/natural_fields.py')
-rw-r--r--worldalign/natural_fields.py224
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()