summaryrefslogtreecommitdiff
path: root/worldalign/synth_cc_battery.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/synth_cc_battery.py')
-rw-r--r--worldalign/synth_cc_battery.py377
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()