"""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()