diff options
Diffstat (limited to 'worldalign/cooccurrence_dictionary.py')
| -rw-r--r-- | worldalign/cooccurrence_dictionary.py | 237 |
1 files changed, 237 insertions, 0 deletions
diff --git a/worldalign/cooccurrence_dictionary.py b/worldalign/cooccurrence_dictionary.py new file mode 100644 index 0000000..4c338e4 --- /dev/null +++ b/worldalign/cooccurrence_dictionary.py @@ -0,0 +1,237 @@ +"""Derive the cross-modal value correspondence from joint structure. + +Marginal frequency pairs values only when the two corpora rank them the +same way, which reporting bias erodes: text mentions the salient, pixels +count the common. Joint structure is sturdier. Which colours co-occur +with which shapes is a property of the world that both modalities +observe, and reporting bias distorts the marginals long before it +scrambles that pattern. + +The correspondence is therefore recovered by matching two small +value-level graphs -- colour-by-shape co-occurrence, estimated per +modality on its own corpus -- with the same spectral-plus-refinement +solver used at scene level. Both corpora stay disjoint; no pair is read. +""" + +from __future__ import annotations + +import argparse +import json +from collections import Counter +from pathlib import Path + +import numpy as np +import torch +from scipy.cluster.hierarchy import fcluster, linkage +from scipy.optimize import linear_sum_assignment +from sklearn.cluster import KMeans +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .synth_set_battery import parse_group_phrases +from .synth_towers import load_image +from .tier0_pipeline import ( + appearance, + extract_objects, + group_objects, +) +from .tier0_dictionary import partition_text_vocabulary, text_token_statistics + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v3") + parser.add_argument("--fit-scenes", type=int, default=3000) + parser.add_argument("--peak-distance", type=int, default=5) + parser.add_argument("--group-threshold", type=float, default=0.35) + parser.add_argument("--shape-classes", type=int, default=6) + parser.add_argument("--restarts", type=int, default=40) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--output", default="artifacts/synth_v3/cooccurrence_dictionary.json" + ) + return parser.parse_args() + + +def text_joint( + captions: list[list[str]], + rows: list[int], + colour_words: list[str], + shape_words: list[str], +) -> np.ndarray: + joint = np.zeros((len(colour_words), len(shape_words))) + colour_index = {word: i for i, word in enumerate(colour_words)} + shape_index = {word: i for i, word in enumerate(shape_words)} + for row in rows: + for phrase in parse_group_phrases(captions[row][0]): + tokens = phrase.split() + colour = next((colour_index[t] for t in tokens if t in colour_index), None) + shape = next( + ( + shape_index[t.rstrip("es") if t.rstrip("es") in shape_index else t] + for t in tokens + if t in shape_index or t.rstrip("es") in shape_index + ), + None, + ) + if colour is not None and shape is not None: + joint[colour, shape] += 1.0 + return joint + + +def vision_joint( + rows: list[int], + image_dir: Path, + args: argparse.Namespace, + colour_classes: int, +) -> tuple[np.ndarray, KMeans, KMeans, dict[int, int]]: + groups, shapes = [], [] + for row in tqdm(rows, desc="vision joint"): + objects = extract_objects( + load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance + ) + members = group_objects(objects, args.group_threshold) + if not members: + continue + features = np.stack([appearance(item) for item in objects]) + if len(objects) == 1: + labels = np.array([0]) + else: + labels = fcluster( + linkage(features, "complete"), args.group_threshold, "distance" + ) + buckets: dict[int, list[dict]] = {} + for item, label in zip(objects, labels): + buckets.setdefault(int(label), []).append(item) + for group, bucket in zip(members, buckets.values()): + groups.append(group) + shapes.append(np.mean([item["shape"] for item in bucket], axis=0)) + colour_model = KMeans(colour_classes, n_init=10, random_state=args.seed).fit( + np.stack([group["rgb"] for group in groups]).astype(np.float64) + ) + shape_model = KMeans(args.shape_classes, n_init=10, random_state=args.seed).fit( + np.stack(shapes).astype(np.float64) + ) + colour_labels = colour_model.labels_ + shape_labels = shape_model.labels_ + frequency = Counter(colour_labels.tolist()) + rank = {label: index for index, (label, _) in enumerate(frequency.most_common())} + joint = np.zeros((colour_classes, args.shape_classes)) + for colour, shape in zip(colour_labels, shape_labels): + joint[rank[int(colour)], int(shape)] += 1.0 + return joint, colour_model, shape_model, rank + + +def normalise_joint(joint: np.ndarray) -> np.ndarray: + """Row-normalised joint: the shape profile of each colour.""" + return joint / joint.sum(axis=1, keepdims=True).clip(1e-9) + + +def match_values( + text: np.ndarray, vision: np.ndarray, restarts: int, seed: int +) -> tuple[np.ndarray, np.ndarray, float]: + """Align colour rows and shape columns of two joint tables. + + Alternating assignment: given a column correspondence, rows are + matched by profile similarity; given rows, columns are rematched. + Iterating from many random column starts and keeping the best + agreement avoids the trivial fixed point. Row indices are returned as + vision-colour to text-colour. + """ + generator = np.random.default_rng(seed) + text_profile = normalise_joint(text) + vision_profile = normalise_joint(vision) + shapes = text.shape[1] + best_row, best_column, best_value = None, None, -np.inf + for restart in range(restarts): + column = ( + np.arange(shapes) if restart == 0 else generator.permutation(shapes) + ) + row = None + for _ in range(20): + # Rows: vision colour i against text colour j under current columns. + score = vision_profile @ text_profile[:, column].T + _, row = linear_sum_assignment(-score) + # Columns: vision shape a against text shape b under current rows. + column_score = vision_profile.T @ text_profile[row] + _, new_column = linear_sum_assignment(-column_score) + if np.array_equal(new_column, column): + break + column = new_column + value = float((vision_profile * text_profile[row][:, column]).sum()) + if value > best_value: + best_row, best_column, best_value = row, column, value + return best_row, best_column, best_value + + +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"] + scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"] + image_dir = Path(manifest["image_dir"]) + + families = partition_text_vocabulary( + text_token_statistics(captions, manifest["text_only_train"]) + ) + colour_words = families["colour_words"] + from .synth_world import SHAPES + + shape_words = list(SHAPES) + text_table = text_joint( + captions, manifest["text_only_train"], colour_words, shape_words + ) + vision_table, colour_model, _, rank = vision_joint( + manifest["vision_only_train"][: args.fit_scenes], + image_dir, + args, + len(colour_words), + ) + + row, column, agreement = match_values( + text_table, vision_table, args.restarts, args.seed + ) + + # Evaluation only: names attached to vision clusters by nearest true RGB. + from .synth_world import COLORS + + names = list(COLORS) + reference = np.asarray([COLORS[name] for name in names], dtype=np.float64) / 255.0 + inverse_rank = {index: label for label, index in rank.items()} + cluster_name = { + index: names[ + int(np.argmin(((colour_model.cluster_centers_[inverse_rank[index]] - reference) ** 2).sum(1))) + ] + for index in range(len(colour_words)) + } + marginal_correct = sum( + 1 + for index, word in enumerate(colour_words) + if index < len(cluster_name) and word == cluster_name[index] + ) + joint_correct = sum( + 1 + for vision_index, text_index in enumerate(row) + if colour_words[text_index] == cluster_name[vision_index] + ) + report = { + "protocol": ( + "Colour-by-shape joint tables are estimated per modality on " + "disjoint corpora and aligned by alternating assignment; " + "truth is read only to count correct entries." + ), + "colour_words": colour_words, + "agreement": agreement, + "marginal_rank_correct": marginal_correct, + "joint_structure_correct": joint_correct, + "classes": len(colour_words), + } + print(json.dumps({k: report[k] for k in + ("marginal_rank_correct", "joint_structure_correct", "classes")})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
