diff options
Diffstat (limited to 'worldalign/synth_cc_battery.py')
| -rw-r--r-- | worldalign/synth_cc_battery.py | 377 |
1 files changed, 377 insertions, 0 deletions
diff --git a/worldalign/synth_cc_battery.py b/worldalign/synth_cc_battery.py new file mode 100644 index 0000000..be28a01 --- /dev/null +++ b/worldalign/synth_cc_battery.py @@ -0,0 +1,377 @@ +"""Upper-bound set battery: connected-component sprites as vision sets. + +The learned towers have not yet produced object states, which leaves two +hypotheses entangled: the set-kernel machinery could be wrong, or only +the towers could be short. This battery separates them. On this world +the background is flat, so connected bright components ARE the objects; +per-component centered sprites are model-free object states of the +minimal-world-knowledge class (like the color anchors on natural data: +declared, fixed, no learning). If set-kernel relation fields built from +these pass the gate and support recovery, the machinery is validated and +the remaining gap is exactly "an SSL objective that discovers objects". + +Text sets are the per-group phrase states of the set battery. Hidden +pairs score orderings only. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import numpy as np +import torch +import torch.nn.functional as F +from scipy import ndimage +from tqdm import tqdm + +from .common import read_json, seed_everything, write_json +from .manifold_gate import standardize_relation +from .ricci_control import run_gates +from .synth_set_battery import ( + parse_group_phrases, + set_similarity_field, + text_group_sets, +) +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt") + parser.add_argument("--split", choices=["val", "test"], default="test") + parser.add_argument("--samples", type=int, default=512) + parser.add_argument("--sprite-window", type=int, default=56) + parser.add_argument("--merge-distance", type=float, default=30.0) + parser.add_argument( + "--features", choices=["pixels", "descriptors", "onehot"], default="descriptors" + ) + parser.add_argument("--text-mode", choices=["lm", "bow"], default="bow") + parser.add_argument( + "--kernel", choices=["matching", "moment"], default="matching", + help="moment: symmetric-tensor (Fock) set kernel, no matching step", + ) + parser.add_argument("--vision-views", type=int, default=4) + parser.add_argument("--random-perms", type=int, default=300) + parser.add_argument("--descent-restarts", type=int, default=5) + parser.add_argument("--descent-max-steps", type=int, default=200000) + parser.add_argument( + "--descent-objective", default="mse", choices=["mse", "m30_total"] + ) + parser.add_argument("--descent-verify-top", type=int, default=64) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260731) + parser.add_argument( + "--output", default="artifacts/synth_v0/cc_battery_gate.json" + ) + return parser.parse_args() + + +def component_sprites( + image: torch.Tensor, window: int, merge_distance: float +) -> tuple[torch.Tensor, torch.Tensor]: + """Centered sprites of bright connected components, ring-merged. + + Same-group objects sit on a small ring; merging nearby components + reassembles groups so a sprite carries multiplicity as pattern. + """ + 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 torch.zeros(1, 3 * window * window), torch.ones(1) + centers = np.array(ndimage.center_of_mass(foreground, labels, range(1, count + 1))) + sizes = ndimage.sum(foreground, labels, range(1, count + 1)) + # Merge components whose centers are close (ring members). + 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) + + height, width = foreground.shape + half = window // 2 + padded = np.pad(array, ((half, half), (half, half), (0, 0))) + padded_mask = np.pad(foreground, half) + sprites, weights = [], [] + for members in groups.values(): + member_mask = np.isin(labels, [m + 1 for m in members]) + mass = float(member_mask.sum()) + ys, xs = np.nonzero(member_mask) + cy, cx = int(ys.mean()), int(xs.mean()) + patch = padded[cy : cy + window, cx : cx + window].copy() + mask_patch = padded_mask[cy : cy + window, cx : cx + window] + patch[~mask_patch] = 0.0 + sprites.append(torch.from_numpy(patch).float().flatten()) + weights.append(mass) + weights = torch.tensor(weights) + return torch.stack(sprites), weights / weights.sum().clamp_min(1e-8) + + +HUE_CENTERS = { + "red": 0.0, "orange": 30.0, "yellow": 60.0, "green": 120.0, + "cyan": 180.0, "blue": 220.0, "purple": 275.0, "pink": 330.0, +} + + +def component_descriptors( + image: torch.Tensor, merge_distance: float +) -> tuple[torch.Tensor, torch.Tensor]: + """Rotation-invariant group descriptors: color classes, log-area, + member multiplicity, and normalized single-member area.""" + import colorsys + + 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 torch.zeros(1, 19), torch.ones(1) + 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) + descriptors, weights = [], [] + total = foreground.size + for members in groups.values(): + member_mask = np.isin(labels, [m + 1 for m in members]) + pixels = array[member_mask] + mean_rgb = pixels.mean(0) + h, s_, v = colorsys.rgb_to_hsv(*mean_rgb.tolist()) + hue = h * 360.0 + vector = np.zeros(16) + if v > 0.75 and s_ < 0.25: + vector[8] = 1.0 # white + elif v < 0.25: + vector[9] = 1.0 # black + elif s_ < 0.25: + vector[10] = 1.0 # gray + elif abs(hue - 30.0) < 25.0 and v < 0.55: + vector[11] = 1.0 # brown + else: + for index, center in enumerate(HUE_CENTERS.values()): + distance = min(abs(hue - center), 360.0 - abs(hue - center)) + if distance < 25.0: + vector[index] = 1.0 + break + member_count = len(members) + area = float(member_mask.sum()) / total + vector[12] = np.log(area + 1e-6) / 6.0 + vector[13] = (member_count - 1) / 3.0 + vector[14] = np.log(area / member_count + 1e-6) / 6.0 + vector[15] = 1.0 + # Rotation-invariant shape features of the largest single member: + # compactness, convexity, and inertia eccentricity separate the + # six shapes without orientation. + largest = max(members, key=lambda m: (labels == m + 1).sum()) + single = labels == largest + 1 + area_px = float(single.sum()) + eroded = ndimage.binary_erosion(single) + perimeter = float((single & ~eroded).sum()) + compactness = 4.0 * np.pi * area_px / max(perimeter, 1.0) ** 2 + ys, xs = np.nonzero(single) + ys = ys - ys.mean(); xs = xs - xs.mean() + cov = np.cov(np.stack([xs, ys])) + 1e-6 * np.eye(2) + eigenvalues = np.linalg.eigvalsh(cov) + eccentricity = float(1.0 - eigenvalues[0] / eigenvalues[1]) + hull_span = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1) + boxfill = area_px / max(hull_span, 1.0) + shape_vector = np.array([compactness, eccentricity, boxfill]) + vector = np.concatenate([vector, shape_vector]) + descriptors.append(torch.tensor(vector, dtype=torch.float32)) + weights.append(float(member_mask.sum())) + weights = torch.tensor(weights) + return torch.stack(descriptors), weights / weights.sum().clamp_min(1e-8) + + +def onehot_descriptors( + raw_sets: list[torch.Tensor], +) -> list[torch.Tensor]: + """Factor-mirrored one-hot recoding of descriptor sets. + + Sizes are binned by corpus terciles of single-member log-area with the + middle bin unmarked, mirroring the text side where medium size has no + word; shapes are k-means clusters of the rotation-invariant shape + features. Both statistics come from the evaluated corpus itself, + unimodally. Output channels mirror the bag-of-words support: color + (12), small/large (2), count (4), shape cluster (6). + """ + from sklearn.cluster import KMeans + + all_groups = torch.cat(raw_sets) + # The renderer shrinks radii with member count, so raw single-member + # area confounds size class with multiplicity; residualize log-area on + # member count before binning (unimodal statistics). + counts_all = (all_groups[:, 13] * 3.0).round().clamp(0, 3) + area_all = all_groups[:, 14] + count_means = {} + for value in (0.0, 1.0, 2.0, 3.0): + chosen = counts_all == value + count_means[value] = float(area_all[chosen].mean()) if chosen.any() else 0.0 + adjusted_all = area_all - torch.tensor( + [count_means[float(v)] for v in counts_all] + ) + low, high = adjusted_all.quantile(1.0 / 3.0), adjusted_all.quantile(2.0 / 3.0) + shape_features = all_groups[:, 16:19].numpy() + clusters = KMeans(n_clusters=6, n_init=10, random_state=0).fit(shape_features) + recoded = [] + for groups in raw_sets: + vectors = torch.zeros(len(groups), 24) + vectors[:, :12] = groups[:, :12] + counts_here = (groups[:, 13] * 3.0).round().clamp(0, 3) + adjusted = groups[:, 14] - torch.tensor( + [count_means[float(v)] for v in counts_here] + ) + vectors[:, 12] = (adjusted <= low).float() # small + vectors[:, 13] = (adjusted >= high).float() # large + counts = (groups[:, 13] * 3.0).round().long().clamp(0, 3) + vectors[torch.arange(len(groups)), 14 + counts] = 1.0 + labels = clusters.predict(groups[:, 16:19].numpy()) + vectors[torch.arange(len(groups)), 18 + labels] = 1.0 + recoded.append(vectors) + return recoded + + +def moment_field(sets: list[torch.Tensor]) -> torch.Tensor: + """Symmetric-tensor (second-quantized) set kernel, matching-free. + + phi(S) concatenates the degree-1 and degree-2 moments of the set; the + field is the Gram matrix of normalized phi. No assignment step, so no + matching-value or set-size bias can enter. + """ + phis = [] + for members in sets: + m1 = members.mean(0) + m2 = (members[:, :, None] * members[:, None, :]).mean(0).flatten() + phi = torch.cat([m1, m2]) + phis.append(phi / phi.norm().clamp_min(1e-9)) + stacked = torch.stack(phis) + return stacked @ stacked.T + + +def phrase_bow_sets( + rows: list[int], captions: list[list[str]], vocabulary: list[str] +) -> list[torch.Tensor]: + index = {word: i for i, word in enumerate(vocabulary)} + sets = [] + for row in rows: + phrases = parse_group_phrases(captions[row][0]) + vectors = torch.zeros(len(phrases), len(index)) + for p, phrase in enumerate(phrases): + for token in phrase.split(): + if token in index: + vectors[p, index[token]] += 1.0 + sets.append(F.normalize(vectors, dim=-1)) + return sets + + +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"] + rows = manifest[args.split][: args.samples] + image_dir = Path(manifest["image_dir"]) + + per_view_fields = [] + set_sizes = [] + for view in range(args.vision_views): + vision_sets = [] + raw_sets = [] + for row in tqdm(rows, desc=f"cc sprites v{view}"): + image = load_image(image_dir / f"scene{row:06d}_v{view}.png") + if args.features in ("descriptors", "onehot"): + sprites, _ = component_descriptors(image, args.merge_distance) + else: + sprites, _ = component_sprites( + image, args.sprite_window, args.merge_distance + ) + raw_sets.append(sprites) + if view == 0: + set_sizes.append(len(sprites)) + if args.features == "onehot": + vision_sets = [ + F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets) + ] + else: + vision_sets = [F.normalize(v, dim=-1) for v in raw_sets] + if args.kernel == "moment": + per_view_fields.append(moment_field(vision_sets)) + else: + per_view_fields.append(set_similarity_field(vision_sets)) + + if args.text_mode == "bow": + text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"]) + else: + text_sets = text_group_sets(rows, captions, args) + if args.kernel == "moment": + text_field = moment_field(text_sets) + else: + text_field = set_similarity_field(text_sets) + # Mass weighting inside the matching corrupts similarity grading + # (0.16 vs 0.45 against soft truth); match unweighted. Averaging the + # per-view fields cancels segmentation errors across resampled layouts. + visual_field = torch.stack(per_view_fields).mean(0) + visual_channels = standardize_relation(visual_field.double())[0][None] + text_channels = standardize_relation(text_field.double())[0][None] + generator = torch.Generator().manual_seed(args.seed) + report = { + "protocol": ( + "Vision sets are model-free connected-component sprites " + "(declared minimal world knowledge); text sets are per-group " + "phrase states. Hidden pairs score orderings only." + ), + "split": args.split, + "samples": len(rows), + "mean_vision_set_size": float(np.mean(set_sizes)), + **run_gates(text_channels, visual_channels, args, generator), + } + verdict = { + "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], + "improving_fraction": report["gate_b"]["improving_fraction"], + "descent_keeps": report["descent_from_true"]["final_accuracy"], + "true_mse": report["gate_a"]["true"]["mse"], + "best_random_descent": min( + (r["final_objective"] for r in report["descent_from_random"]), + default=None, + ), + } + verdict["counterfeit_found"] = bool( + verdict["best_random_descent"] is not None + and verdict["best_random_descent"] < verdict["true_mse"] + ) + report["verdict"] = verdict + print(json.dumps({"verdict": verdict})) + write_json(args.output, report) + print(f"Wrote {args.output}") + + +if __name__ == "__main__": + main() |
