"""Tier 0 pipeline: unpaired corpora to aligned relation fields. Consolidates the components that produced the synthetic world's end-to-end result, each of which was developed and measured separately: - watershed object extraction, which splits touching ring members that connected components merge (exact object count 74.2% to 100%); - Gestalt appearance grouping, which assembles members into groups by shared colour, size, and shape rather than by a distance threshold (exact group count 71.9% to 90.2%); - size classes from radial extent under one-dimensional k-means per member count, because area confounds size with shape and the classes are gap-separated rather than equally populated (52.7% to 87.9%); - text factor families from mutual exclusivity within a group phrase; - the cross-modal value correspondence from marginal frequency rank. Nothing crosses modalities except the frequency ranking, and the two corpora it reads are disjoint: no instance appears on both sides. """ from __future__ import annotations import argparse import json from collections import Counter, defaultdict from pathlib import Path import numpy as np import torch import torch.nn.functional as F from scipy import ndimage from scipy.cluster.hierarchy import fcluster, linkage from sklearn.cluster import KMeans from tqdm import tqdm from .common import read_json, seed_everything, write_json from .synth_cc_battery import moment_field from .synth_set_battery import parse_group_phrases from .synth_towers import load_image from .tier0_dictionary import partition_text_vocabulary, text_token_statistics NUMBER_TO_COUNT = {"two": 2, "three": 3, "four": 4} SINGULAR = ("a", "an") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--data-dir", default="artifacts/synth_v1") parser.add_argument("--split", choices=["val", "test"], default="test") parser.add_argument("--samples", type=int, default=256) parser.add_argument("--offset", type=int, default=0) parser.add_argument("--fit-scenes", type=int, default=1500) parser.add_argument("--peak-distance", type=int, default=5) parser.add_argument("--group-threshold", type=float, default=0.35) parser.add_argument("--views", type=int, default=0, help="0 uses manifest.") parser.add_argument("--seed", type=int, default=0) parser.add_argument("--output", required=True) parser.add_argument("--states-output", default="") parser.add_argument( "--text-omit", default="", help="Comma-separated factors the captions never state " "(colour, count, size). Vision still sees them.", ) return parser.parse_args() def extract_objects(image: torch.Tensor, peak_distance: int) -> list[dict]: """Foreground objects, splitting touching ones by distance watershed.""" 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 distance = ndimage.distance_transform_edt(foreground) window = 2 * peak_distance + 1 peaks = ( distance >= ndimage.maximum_filter(distance, size=window) - 1e-9 ) & (distance > 1.0) labels, count = ndimage.label(peaks) if count == 0: labels, count = ndimage.label(foreground) else: while True: grown = ndimage.grey_dilation(labels, size=3) take = (labels == 0) & foreground & (grown > 0) if not take.any(): break labels = np.where(take, grown, labels) count = labels.max() objects = [] for index in range(1, count + 1): mask = labels == index area = float(mask.sum()) if area < 6: continue ys, xs = np.nonzero(mask) centred_y, centred_x = ys - ys.mean(), xs - xs.mean() covariance = np.cov(np.stack([centred_x, centred_y])) + 1e-6 * np.eye(2) eigenvalues = np.linalg.eigvalsh(covariance) eroded = ndimage.binary_erosion(mask) perimeter = max(float((mask & ~eroded).sum()), 1.0) box = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1) objects.append( { "rgb": array[mask].mean(0), "area": area, "extent": float(np.hypot(centred_y, centred_x).max()), "shape": np.array( [ 4 * np.pi * area / perimeter**2, 1 - eigenvalues[0] / eigenvalues[1], area / max(box, 1), ] ), } ) return objects def appearance(item: dict) -> np.ndarray: return np.concatenate( [item["rgb"] * 3.0, [np.log(item["area"] + 1e-6) * 0.6], item["shape"]] ) def group_objects(objects: list[dict], threshold: float) -> list[dict]: """Members of one group share appearance; group by similarity.""" if not objects: return [] if len(objects) == 1: labels = np.array([0]) else: features = np.stack([appearance(item) for item in objects]) labels = fcluster(linkage(features, "complete"), threshold, "distance") buckets: dict[int, list[dict]] = {} for item, label in zip(objects, labels): buckets.setdefault(int(label), []).append(item) return [ { "rgb": np.mean([item["rgb"] for item in members], axis=0), "members": len(members), "extent": float(np.mean([item["extent"] for item in members])), } for members in buckets.values() ] class VisionCoder: """Colour classes and size classes fitted on the vision corpus alone.""" def __init__(self, groups: list[dict], classes: int, seed: int) -> None: colours = np.stack([group["rgb"] for group in groups]).astype(np.float64) self.colour_model = KMeans(classes, n_init=10, random_state=seed).fit(colours) frequency = Counter(self.colour_model.labels_.tolist()) self.colour_rank = { label: rank for rank, (label, _) in enumerate(frequency.most_common()) } by_count: dict[int, list[float]] = defaultdict(list) for group in groups: by_count[min(group["members"], 4)].append(group["extent"]) self.size_models = {} for count, extents in by_count.items(): model = KMeans(3, n_init=10, random_state=seed).fit( np.asarray(extents, dtype=np.float64)[:, None] ) order = np.argsort(model.cluster_centers_[:, 0]) self.size_models[count] = ( model, {int(label): rank for rank, label in enumerate(order)}, ) def encode(self, groups: list[dict], classes: int) -> torch.Tensor: if not groups: groups = [{"rgb": np.zeros(3), "members": 1, "extent": 1.0}] colours = self.colour_model.predict( np.stack([group["rgb"] for group in groups]).astype(np.float64) ) vectors = [] for group, colour in zip(groups, colours): model, order = self.size_models[min(group["members"], 4)] size = order[ int(model.predict(np.array([[group["extent"]]], dtype=np.float64))[0]) ] vectors.append( factor_vector(self.colour_rank[int(colour)], group["members"], size, classes) ) return F.normalize(torch.stack(vectors), dim=-1) def factor_vector(colour: int, count: int, size: int, classes: int) -> torch.Tensor: vector = torch.zeros(classes + 4 + 3) vector[colour] = 1.0 vector[classes + min(count - 1, 3)] = 1.0 vector[classes + 4 + size] = 1.0 return vector def encode_caption( caption: str, colour_rank: dict[str, int], classes: int, omit: frozenset = frozenset() ) -> torch.Tensor: """Encode a caption, optionally suppressing factors the text never states. Suppressing a factor makes the caption describe strictly less of the scene while vision continues to see all of it, which is the controlled form of the situation on photographs: a region description and a patch descriptor overlap on some aspects and not others. It is the intervention that tests whether corpus overlap is what sets the width of the shared spectrum. """ vectors = [] for phrase in parse_group_phrases(caption): tokens = phrase.split() colour = next((colour_rank[t] for t in tokens if t in colour_rank), 0) count = ( 1 if any(token in SINGULAR for token in tokens) else next( (NUMBER_TO_COUNT[t] for t in tokens if t in NUMBER_TO_COUNT), 1 ) ) size = 0 if "small" in tokens else (2 if "large" in tokens else 1) vector = factor_vector(colour, count, size, classes) if "colour" in omit: vector[:classes] = 0.0 if "count" in omit: vector[classes : classes + 4] = 0.0 if "size" in omit: vector[classes + 4 :] = 0.0 vectors.append(vector) if not vectors: vectors = [factor_vector(0, 1, 1, classes)] stacked = torch.stack(vectors) if stacked.abs().sum() == 0: stacked = stacked + 1.0 return F.normalize(stacked, dim=-1) def moment_state(states: torch.Tensor) -> torch.Tensor: first = states.mean(0) second = (states[:, :, None] * states[:, None, :]).mean(0).flatten() return torch.cat([first, second]) 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"] image_dir = Path(manifest["image_dir"]) views = args.views or manifest["visual_views"] rows = manifest[args.split][args.offset : args.offset + args.samples] families = partition_text_vocabulary( text_token_statistics(captions, manifest["text_only_train"]) ) colour_words = families["colour_words"] classes = len(colour_words) colour_rank = {word: rank for rank, word in enumerate(colour_words)} fit_groups = [ group for row in tqdm( manifest["vision_only_train"][: args.fit_scenes], desc="fit vision" ) for group in group_objects( extract_objects( load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance ), args.group_threshold, ) ] coder = VisionCoder(fit_groups, classes, args.seed) per_view_fields = [] view_states = [] for view in range(views): sets = [ coder.encode( group_objects( extract_objects( load_image(image_dir / f"scene{row:06d}_v{view}.png"), args.peak_distance, ), args.group_threshold, ), classes, ) for row in tqdm(rows, desc=f"encode view {view}") ] per_view_fields.append(moment_field(sets)) if view == 0: view_states = [moment_state(item) for item in sets] visual_field = torch.stack(per_view_fields).mean(0) omit = frozenset(f for f in args.text_omit.split(",") if f) text_sets = [ encode_caption(captions[row][0], colour_rank, classes, omit) for row in rows ] text_field = moment_field(text_sets) mask = ~np.eye(len(rows), 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, "rows": rows}, args.output, ) if args.states_output: torch.save( { "vision_states": torch.stack(view_states), "text_states": torch.stack([moment_state(item) for item in text_sets]), "rows": rows, }, args.states_output, ) summary = { "data_dir": args.data_dir, "split": args.split, "samples": len(rows), "colour_classes": classes, "text_families": { key: families[key] for key in ("count_words", "colour_words", "size_words") }, "field_correlation_at_truth": correlation, "note": ( "The dictionary is derived from disjoint corpora; the " "correlation is a diagnostic computed with hidden pairs and " "never used by the pipeline." ), } print(json.dumps({"field_correlation_at_truth": correlation})) write_json(args.output.replace(".pt", ".json"), summary) print(f"Wrote {args.output}") if __name__ == "__main__": main()