summaryrefslogtreecommitdiff
path: root/worldalign/tier0_dictionary.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/tier0_dictionary.py')
-rw-r--r--worldalign/tier0_dictionary.py324
1 files changed, 324 insertions, 0 deletions
diff --git a/worldalign/tier0_dictionary.py b/worldalign/tier0_dictionary.py
new file mode 100644
index 0000000..96fcbac
--- /dev/null
+++ b/worldalign/tier0_dictionary.py
@@ -0,0 +1,324 @@
+"""Tier 0: derive the cross-modal value correspondence, never declare it.
+
+A declared lexicon ("red" means hue 0) is a hand-supplied cross-modal
+prior. Tier 0 forbids it. Every factor value correspondence is instead
+recovered from unimodal statistics of the two disjoint training splits:
+
+- ordered factors (count, size) match by their intrinsic order, with the
+ small residual ambiguity enumerated and settled by the alignment
+ criterion rather than by assertion;
+- unordered factors (colour, shape) match by marginal frequency rank,
+ which is a unimodal observable on both sides. This works exactly when
+ the world's factor marginals are non-uniform -- true of real corpora
+ and of the skewed synthetic world, false of the uniform one, where the
+ correspondence is information-theoretically unrecoverable.
+
+The recovered dictionary is then applied to build comparable object
+descriptors. Hidden pairs are used only to report how many entries the
+derivation got right; nothing here reads them.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from collections import Counter
+from pathlib import Path
+
+import numpy as np
+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
+from .synth_set_battery import parse_group_phrases
+from .synth_towers import load_image
+
+SINGULAR_ARTICLES = ("a", "an")
+STOPWORDS = {
+ "there", "are", "the", "picture", "shows", "you", "can", "see",
+ "is", "of", "left", "right", "above", "below", "and",
+}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v1")
+ parser.add_argument("--text-scenes", type=int, default=6000)
+ parser.add_argument("--vision-scenes", type=int, default=3000)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--colour-classes", type=int, default=10)
+ parser.add_argument("--shape-classes", type=int, default=6)
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v1/tier0_dictionary.json"
+ )
+ return parser.parse_args()
+
+
+def text_token_statistics(
+ captions: list[list[str]], rows: list[int]
+) -> dict:
+ """Group-slot token frequencies and singular/plural association.
+
+ Everything here reads captions only. Number words are identified by
+ their association with singular noun forms and their mutual exclusion
+ inside a group phrase; the remaining modifier vocabulary splits into
+ frequency-ranked classes.
+ """
+ phrase_total = 0
+ slot_counts: Counter = Counter()
+ with_singular: Counter = Counter()
+ total_singular = 0
+ cooccurrence: Counter = Counter()
+ for row in rows:
+ for phrase in parse_group_phrases(captions[row][0]):
+ tokens = [
+ token
+ for token in re.findall(r"[a-z]+|\d+", phrase.lower())
+ if token not in STOPWORDS
+ ]
+ if not tokens:
+ continue
+ phrase_total += 1
+ head = tokens[-1]
+ singular = not head.endswith("s")
+ total_singular += singular
+ modifiers = tokens[:-1]
+ for token in modifiers:
+ slot_counts[token] += 1
+ with_singular[token] += singular
+ for first in modifiers:
+ for second in modifiers:
+ if first != second:
+ cooccurrence[(first, second)] += 1
+ return {
+ "slot_counts": slot_counts,
+ "with_singular": with_singular,
+ "cooccurrence": cooccurrence,
+ "total_singular": total_singular,
+ "phrase_total": phrase_total,
+ }
+
+
+def partition_text_vocabulary(stats: dict) -> dict:
+ """Discover modifier families by mutual exclusivity, then name them.
+
+ Tokens of one factor never co-occur inside a group phrase, so families
+ are maximal mutually exclusive sets: greedily place each token in the
+ first family none of whose members it ever co-occurs with. Families are
+ then identified by two unimodal signals -- coverage (how many phrases
+ carry a member) and morphological association (whether the choice
+ predicts the head noun's plural suffix). No token list is declared.
+ """
+ counts = stats["slot_counts"]
+ cooccurrence = stats["cooccurrence"]
+ phrases = stats["phrase_total"]
+ ordered = sorted(counts, key=lambda token: -counts[token])
+ families: list[list[str]] = []
+ for token in ordered:
+ for family in families:
+ if all(
+ cooccurrence[(token, member)] == 0
+ and cooccurrence[(member, token)] == 0
+ for member in family
+ ):
+ family.append(token)
+ break
+ else:
+ families.append([token])
+ described = []
+ for family in families:
+ coverage = sum(counts[token] for token in family) / max(phrases, 1)
+ rates = [
+ stats["with_singular"][token] / max(counts[token], 1)
+ for token in family
+ ]
+ described.append(
+ {
+ "tokens": sorted(family, key=lambda token: -counts[token]),
+ "coverage": coverage,
+ "morphology_spread": float(max(rates) - min(rates)),
+ }
+ )
+ described.sort(key=lambda item: -item["coverage"])
+ # The count family is the near-complete family whose choice predicts the
+ # plural suffix; the other near-complete family is the dominant
+ # unordered attribute; partial families are optional modifiers.
+ complete = [item for item in described if item["coverage"] > 0.8]
+ partial = [item for item in described if item["coverage"] <= 0.8]
+ complete.sort(key=lambda item: -item["morphology_spread"])
+ count_family = complete[0]["tokens"] if complete else []
+ attribute_families = [item["tokens"] for item in complete[1:]]
+ return {
+ "count_words": count_family,
+ "colour_words": attribute_families[0] if attribute_families else [],
+ "other_attribute_words": attribute_families[1:],
+ "size_words": [item["tokens"] for item in partial],
+ "families_detail": described,
+ }
+
+
+def component_raw(image: torch.Tensor, merge_distance: float) -> dict:
+ """Connected-component groups with raw appearance, no colour rules.
+
+ Returns mean RGB, area fraction, and member count per group. Nothing
+ here quantises colour, so no declared hue boundary enters Tier 0.
+ """
+ from scipy import ndimage
+
+ array = image.permute(1, 2, 0).numpy()
+ background = np.median(array.reshape(-1, 3), axis=0)
+ foreground = np.abs(array - background).sum(-1) > 0.12
+ labels, count = ndimage.label(foreground)
+ if count == 0:
+ return {"rgb": np.zeros((0, 3)), "area": np.zeros(0), "members": np.zeros(0)}
+ centers = np.array(
+ ndimage.center_of_mass(foreground, labels, range(1, count + 1))
+ )
+ parent = list(range(count))
+
+ def find(a: int) -> int:
+ while parent[a] != a:
+ parent[a] = parent[parent[a]]
+ a = parent[a]
+ return a
+
+ for a in range(count):
+ for b in range(a + 1, count):
+ if np.linalg.norm(centers[a] - centers[b]) < merge_distance:
+ parent[find(a)] = find(b)
+ groups: dict[int, list[int]] = {}
+ for a in range(count):
+ groups.setdefault(find(a), []).append(a)
+ rgb, area, members = [], [], []
+ total = foreground.size
+ for group in groups.values():
+ mask = np.isin(labels, [m + 1 for m in group])
+ rgb.append(array[mask].mean(0))
+ area.append(float(mask.sum()) / total)
+ members.append(len(group))
+ return {
+ "rgb": np.stack(rgb),
+ "area": np.array(area),
+ "members": np.array(members, dtype=np.int64),
+ }
+
+
+def vision_value_statistics(
+ rows: list[int], manifest: dict, args: argparse.Namespace
+) -> dict:
+ """Colour-class and size-class frequencies from pixels alone.
+
+ Object colours are clustered in hue-saturation-value space with the
+ requested number of classes; class identity is arbitrary, only the
+ frequency ranking is used downstream.
+ """
+ from sklearn.cluster import KMeans
+
+ image_dir = Path(manifest["image_dir"])
+ raw = [
+ component_raw(
+ load_image(image_dir / f"scene{row:06d}_v0.png"), args.merge_distance
+ )
+ for row in tqdm(rows, desc="vision values")
+ ]
+ rgb = np.concatenate([item["rgb"] for item in raw])
+ # Cluster raw appearance: class boundaries come from the data, not from
+ # a declared hue table.
+ clusters = KMeans(
+ n_clusters=args.colour_classes, n_init=10, random_state=args.seed
+ ).fit(rgb)
+ labels = clusters.labels_
+ return {
+ "colour_frequency": Counter(labels.tolist()),
+ "colour_labels": labels,
+ "cluster_centres": clusters.cluster_centers_,
+ "raw": raw,
+ "clusters": clusters,
+ }
+
+
+def frequency_rank_map(
+ text_words: list[str], vision_frequency: Counter, classes: int
+) -> dict:
+ """Match unordered values by descending marginal frequency."""
+ vision_ranked = [
+ label for label, _ in vision_frequency.most_common(classes)
+ ]
+ return {
+ word: vision_ranked[index]
+ for index, word in enumerate(text_words[:classes])
+ }
+
+
+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"]
+
+ text_rows = manifest["text_only_train"][: args.text_scenes]
+ stats = text_token_statistics(captions, text_rows)
+ families = partition_text_vocabulary(stats)
+
+ vision_rows = manifest["vision_only_train"][: args.vision_scenes]
+ vision = vision_value_statistics(vision_rows, manifest, args)
+
+ dictionary = frequency_rank_map(
+ families["colour_words"], vision["colour_frequency"], args.colour_classes
+ )
+
+ report = {
+ "protocol": (
+ "Factor families and their value correspondence are derived "
+ "from unimodal statistics of disjoint splits: noun-number "
+ "association separates count words, coverage separates colour "
+ "from size words, and marginal frequency rank pairs colour "
+ "values across modalities. No declared lexicon."
+ ),
+ "text_families": {
+ key: families[key]
+ for key in ("count_words", "colour_words", "size_words")
+ },
+ "families_detail": families["families_detail"],
+ "vision_colour_frequency": [
+ [int(label), int(count)]
+ for label, count in vision["colour_frequency"].most_common()
+ ],
+ "derived_colour_map": {
+ word: int(label) for word, label in dictionary.items()
+ },
+ }
+
+ # Evaluation only: how many derived entries are semantically right?
+ scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"]
+ truth_frequency = Counter(
+ group["color"] for row in vision_rows for group in scenes[row]["groups"]
+ )
+ truth_rank = [name for name, _ in truth_frequency.most_common()]
+ text_frequency = Counter(
+ group["color"] for row in text_rows for group in scenes[row]["groups"]
+ )
+ text_rank = [name for name, _ in text_frequency.most_common()]
+ report["evaluation_only"] = {
+ "vision_side_truth_frequency_rank": truth_rank,
+ "text_side_truth_frequency_rank": text_rank,
+ "rank_agreement": float(
+ np.mean([a == b for a, b in zip(truth_rank, text_rank)])
+ ),
+ "text_colour_words_recovered": sorted(
+ set(families["colour_words"][: args.colour_classes])
+ & set(truth_rank)
+ ),
+ }
+ print(json.dumps(report["text_families"], indent=1))
+ print(json.dumps(report["evaluation_only"], indent=1))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()