summaryrefslogtreecommitdiff
path: root/worldalign/synth_world.py
diff options
context:
space:
mode:
authorYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
committerYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
commita62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch)
treeee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign/synth_world.py
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built within each modality and are invariant to how each side labels its own features; the cross-modal bridge is a coupling searched under an energy that is a closed-form functional of one matrix; solving is spectral initialisation followed by exact local refinement. Evidence: in a procedurally generated closed world, blind recovery of a hidden image-caption correspondence reaches 95.3% at 256 scenes against 0.39% chance, and the recovered pairs transfer to 200 held-out scenes at 93.0% exact retrieval with random-pair and shuffled-image controls at or near chance. Cross-modal value correspondence is derived from disjoint corpora rather than declared. On Visual Genome the field correlation reaches 0.656 against the 0.9 that polynomial recovery needs, with the deficit attributed away from segmentation and discretisation. Protocol: no image-text pair enters any objective, optimiser, initialisation, or model selection; hidden pairs score orderings only. Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/synth_world.py')
-rw-r--r--worldalign/synth_world.py506
1 files changed, 506 insertions, 0 deletions
diff --git a/worldalign/synth_world.py b/worldalign/synth_world.py
new file mode 100644
index 0000000..8875fde
--- /dev/null
+++ b/worldalign/synth_world.py
@@ -0,0 +1,506 @@
+"""Synthetic closed world v0: procedural scenes with exact closure.
+
+A scene's discrete state is a set of object groups (count, size, color,
+shape) plus a set of spatial relation constraints between groups. Captions
+mention exactly the discrete state; renders realize it with continuous
+nuisance (positions, jitter) resampled per view. Shared content and
+modality-private variation are therefore separated by construction, and
+every dial -- vocabulary, ontology size, groups per scene, scene count,
+orbit multiplicity, intervention density -- is a generator argument.
+
+Outputs mirror the Flickr pipeline layout: a manifest with disjoint
+vision-only and text-only scene rows plus held-out val/test, an image
+directory with `sceneNNNNNN_vK.png` views, and per-scene caption lists.
+Intervention variants with edit metadata are stored for the response
+battery; nothing downstream reads them yet.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import math
+import random
+from pathlib import Path
+
+from PIL import Image, ImageDraw
+
+from .common import write_json
+
+COLORS = {
+ "red": (205, 49, 49),
+ "orange": (224, 133, 44),
+ "yellow": (229, 213, 74),
+ "green": (64, 168, 75),
+ "blue": (59, 104, 214),
+ "purple": (139, 72, 190),
+ "pink": (228, 136, 179),
+ "white": (238, 238, 238),
+ "gray": (140, 140, 140),
+ "brown": (125, 84, 48),
+}
+SHAPES = ("circle", "square", "triangle", "star", "diamond", "cross")
+SIZES = {"small": (7, 10), "medium": (13, 17), "large": (21, 26)}
+COUNT_WORDS = {1: "one", 2: "two", 3: "three", 4: "four"}
+RELATIONS = ("left of", "right of", "above", "below")
+BACKGROUND = (24, 24, 28)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--output-dir", default="artifacts/synth_v0")
+ parser.add_argument("--image-size", type=int, default=128)
+ parser.add_argument("--vision-only", type=int, default=12000)
+ parser.add_argument("--text-only", type=int, default=12000)
+ parser.add_argument("--val", type=int, default=1000)
+ parser.add_argument("--test", type=int, default=1000)
+ parser.add_argument("--min-groups", type=int, default=2)
+ parser.add_argument("--max-groups", type=int, default=4)
+ parser.add_argument("--visual-views", type=int, default=4)
+ parser.add_argument("--captions", type=int, default=4)
+ parser.add_argument("--interventions", type=int, default=2,
+ help="Edited variants per eval scene.")
+ parser.add_argument("--seed", type=int, default=20260730)
+ parser.add_argument(
+ "--skew",
+ type=float,
+ default=0.0,
+ help="Zipf exponent for colour and shape sampling. Non-uniform "
+ "marginals make the cross-modal value correspondence recoverable "
+ "from frequency alone, with no declared dictionary.",
+ )
+ parser.add_argument(
+ "--correlate",
+ type=float,
+ default=0.0,
+ help="Colour-shape coupling strength. Real worlds pair values "
+ "(bananas are yellow), which makes co-occurrence structure "
+ "informative where bare marginals are not.",
+ )
+ parser.add_argument(
+ "--report-bias",
+ type=float,
+ default=0.0,
+ help="Reporting bias strength. Captions mention rarer colours "
+ "preferentially, so text frequency stops tracking pixel "
+ "frequency -- the situation in real corpora, where nobody "
+ "describes the grey road.",
+ )
+ parser.add_argument(
+ "--texture",
+ action="store_true",
+ help="Dot-textured object fills: within-object structure that "
+ "gives binding objectives a purchase, without changing the "
+ "discrete state or the captions.",
+ )
+ return parser.parse_args()
+
+
+_PROFILE_CACHE: dict[tuple[str, float], list[float]] = {}
+
+
+def colour_shape_profile(colour: str, concentration: float) -> list[float]:
+ """Fixed per-colour shape distribution, deterministic in the colour."""
+ key = (colour, concentration)
+ if key not in _PROFILE_CACHE:
+ local = random.Random(hash(key) & 0xFFFFFFFF)
+ raw = [local.gammavariate(1.0 / concentration, 1.0) for _ in SHAPES]
+ total = sum(raw) or 1.0
+ _PROFILE_CACHE[key] = [value / total for value in raw]
+ return _PROFILE_CACHE[key]
+
+
+def zipf_weights(count: int, exponent: float) -> list[float]:
+ raw = [1.0 / (rank + 1) ** exponent for rank in range(count)]
+ total = sum(raw)
+ return [value / total for value in raw]
+
+
+def sample_group(
+ rng: random.Random, skew: float = 0.0, correlate: float = 0.0
+) -> dict:
+ if skew > 0.0:
+ colors = rng.choices(
+ list(COLORS), weights=zipf_weights(len(COLORS), skew)
+ )[0]
+ else:
+ colors = rng.choice(list(COLORS))
+ if correlate > 0.0:
+ # Each colour carries its own shape distribution, drawn once from a
+ # Dirichlet with concentration 1/correlate. Distinct profiles make
+ # the joint table identifying, as in a real world where objects of
+ # a kind take characteristic forms.
+ weights = colour_shape_profile(colors, correlate)
+ if skew > 0.0:
+ base = zipf_weights(len(SHAPES), skew)
+ weights = [w * b for w, b in zip(weights, base)]
+ shapes = rng.choices(SHAPES, weights=weights)[0]
+ elif skew > 0.0:
+ shapes = rng.choices(
+ SHAPES, weights=zipf_weights(len(SHAPES), skew)
+ )[0]
+ else:
+ shapes = rng.choice(SHAPES)
+ return {
+ "count": rng.randint(1, 4),
+ "size": rng.choice(list(SIZES)),
+ "color": colors,
+ "shape": shapes,
+ }
+
+
+def sample_scene(rng: random.Random, args: argparse.Namespace) -> dict:
+ skew = getattr(args, "skew", 0.0)
+ correlate = getattr(args, "correlate", 0.0)
+ groups = [
+ sample_group(rng, skew, correlate)
+ for _ in range(rng.randint(args.min_groups, args.max_groups))
+ ]
+ # Reject referential ambiguity: color-shape pairs are unique per scene,
+ # so every relational mention has exactly one referent.
+ signatures = [(group["color"], group["shape"]) for group in groups]
+ while len(set(signatures)) < len(signatures):
+ groups = [sample_group(rng, skew, correlate) for _ in range(len(groups))]
+ signatures = [(group["color"], group["shape"]) for group in groups]
+ relation_count = rng.randint(1, min(3, len(groups) * (len(groups) - 1) // 2))
+ pairs = [(a, b) for a in range(len(groups)) for b in range(len(groups)) if a < b]
+ rng.shuffle(pairs)
+ relations = [
+ {"a": a, "b": b, "relation": rng.choice(RELATIONS)}
+ for a, b in pairs[:relation_count]
+ ]
+ return {"groups": groups, "relations": relations}
+
+
+def relation_holds(relation: str, pa: tuple[float, float], pb: tuple[float, float], margin: float) -> bool:
+ if relation == "left of":
+ return pa[0] < pb[0] - margin
+ if relation == "right of":
+ return pa[0] > pb[0] + margin
+ if relation == "above":
+ return pa[1] < pb[1] - margin
+ return pa[1] > pb[1] + margin
+
+
+def group_geometry(group: dict, rng: random.Random, size: int) -> dict:
+ low, high = SIZES[group["size"]]
+ shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1]
+ radius = rng.uniform(low, high) * size / 128.0 * shrink
+ count = group["count"]
+ if count == 1:
+ offsets = [(0.0, 0.0)]
+ extent = radius
+ else:
+ # Ring placement guarantees exact visible multiplicity: adjacent
+ # spacing 2 s sin(pi/k) stays above 2.15 r by construction.
+ ring = 1.10 * 1.075 * radius / math.sin(math.pi / count)
+ phase = rng.uniform(0.0, 2.0 * math.pi)
+ offsets = [
+ (
+ ring * math.cos(phase + 2.0 * math.pi * k / count),
+ ring * math.sin(phase + 2.0 * math.pi * k / count),
+ )
+ for k in range(count)
+ ]
+ extent = ring + radius
+ return {"radius": radius, "offsets": offsets, "extent": extent}
+
+
+def worst_case_extent(group: dict, size: int) -> float:
+ high = SIZES[group["size"]][1] * size / 128.0
+ shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1]
+ radius = high * shrink
+ if group["count"] == 1:
+ return radius
+ return 1.10 * 1.075 * radius / math.sin(math.pi / group["count"]) + radius
+
+
+def place_groups(
+ scene: dict,
+ rng: random.Random,
+ size: int,
+ extents: list[float] | None = None,
+) -> list[tuple[float, float]] | None:
+ margin = size * 0.08
+ if extents is None:
+ extents = [worst_case_extent(group, size) for group in scene["groups"]]
+ for _ in range(300):
+ centers = []
+ feasible = True
+ for extent in extents:
+ low, high = extent + 2.0, size - extent - 2.0
+ if low >= high:
+ feasible = False
+ break
+ centers.append((rng.uniform(low, high), rng.uniform(low, high)))
+ if not feasible:
+ return None
+ if any(
+ math.dist(centers[a], centers[b]) < extents[a] + extents[b] + 4.0
+ for a in range(len(centers))
+ for b in range(a + 1, len(centers))
+ ):
+ continue
+ if all(
+ relation_holds(r["relation"], centers[r["a"]], centers[r["b"]], margin)
+ for r in scene["relations"]
+ ):
+ return centers
+ return None
+
+
+def draw_shape(draw: ImageDraw.ImageDraw, shape: str, x: float, y: float, radius: float, fill: tuple) -> None:
+ if shape == "circle":
+ draw.ellipse([x - radius, y - radius, x + radius, y + radius], fill=fill)
+ elif shape == "square":
+ draw.rectangle([x - radius, y - radius, x + radius, y + radius], fill=fill)
+ elif shape == "triangle":
+ draw.polygon(
+ [(x, y - radius), (x - radius, y + radius), (x + radius, y + radius)],
+ fill=fill,
+ )
+ elif shape == "diamond":
+ draw.polygon(
+ [(x, y - radius), (x + radius, y), (x, y + radius), (x - radius, y)],
+ fill=fill,
+ )
+ elif shape == "cross":
+ arm = radius * 0.42
+ draw.rectangle([x - arm, y - radius, x + arm, y + radius], fill=fill)
+ draw.rectangle([x - radius, y - arm, x + radius, y + arm], fill=fill)
+ else: # star
+ points = []
+ for k in range(10):
+ r = radius if k % 2 == 0 else radius * 0.45
+ angle = -math.pi / 2 + k * math.pi / 5
+ points.append((x + r * math.cos(angle), y + r * math.sin(angle)))
+ draw.polygon(points, fill=fill)
+
+
+def render_scene(scene: dict, rng: random.Random, size: int, texture: bool = False) -> Image.Image | None:
+ geometries = [group_geometry(group, rng, size) for group in scene["groups"]]
+ centers = place_groups(
+ scene, rng, size, extents=[g["extent"] for g in geometries]
+ )
+ if centers is None:
+ return None
+ shade = rng.randint(-6, 6)
+ image = Image.new("RGB", (size, size), tuple(c + shade for c in BACKGROUND))
+ draw = ImageDraw.Draw(image)
+ for group, geometry, center in zip(scene["groups"], geometries, centers):
+ base = tuple(
+ min(255, max(0, channel + rng.randint(-10, 10)))
+ for channel in COLORS[group["color"]]
+ )
+ for off in geometry["offsets"]:
+ draw_shape(
+ draw,
+ group["shape"],
+ center[0] + off[0],
+ center[1] + off[1],
+ geometry["radius"],
+ base,
+ )
+ return image
+
+
+def group_phrase(group: dict, rng: random.Random) -> str:
+ size_word = "" if group["size"] == "medium" else group["size"] + " "
+ plural = "es" if group["shape"] == "cross" else "s"
+ noun = group["shape"] + (plural if group["count"] > 1 else "")
+ count_word = COUNT_WORDS[group["count"]] if group["count"] > 1 else (
+ "a" if size_word == "" or size_word[0] not in "aeiou" else "an"
+ )
+ return f"{count_word} {size_word}{group['color']} {noun}"
+
+
+def caption_scene(
+ scene: dict, rng: random.Random, report_bias: float = 0.0
+) -> str:
+ order = list(range(len(scene["groups"])))
+ rng.shuffle(order)
+ if report_bias > 0.0 and len(order) > 1:
+ # Mention probability falls with the colour's population frequency,
+ # so the text marginal inverts the pixel marginal.
+ ranks = {name: index for index, name in enumerate(COLORS)}
+ keep = []
+ for index in order:
+ rank = ranks[scene["groups"][index]["color"]]
+ salience = ((rank + 1) / len(COLORS)) ** report_bias
+ if rng.random() < 0.25 + 0.75 * salience:
+ keep.append(index)
+ order = keep or order[:1]
+ phrases = [group_phrase(scene["groups"][index], rng) for index in order]
+ if len(phrases) > 1:
+ listed = ", ".join(phrases[:-1]) + " and " + phrases[-1]
+ else:
+ listed = phrases[0]
+ opener = rng.choice(["there are", "the picture shows", "you can see"])
+ sentences = [f"{opener} {listed}."]
+ relations = list(scene["relations"])
+ rng.shuffle(relations)
+ mentioned = set(order)
+ relations = [
+ relation
+ for relation in relations
+ if relation["a"] in mentioned and relation["b"] in mentioned
+ ]
+ for relation in relations:
+ subject = group_phrase(scene["groups"][relation["a"]], rng)
+ target = group_phrase(scene["groups"][relation["b"]], rng)
+ sentences.append(f"the {subject.split(' ', 1)[1]} {'are' if scene['groups'][relation['a']]['count'] > 1 else 'is'} {relation['relation']} the {target.split(' ', 1)[1]}.")
+ return " ".join(sentences)
+
+
+def intervene(scene: dict, rng: random.Random) -> tuple[dict, dict]:
+ edited = json.loads(json.dumps(scene))
+ kinds = ["recolor", "count", "remove", "relation"]
+ if len(edited["groups"]) <= 2:
+ kinds.remove("remove")
+ if not edited["relations"]:
+ kinds.remove("relation")
+ kind = rng.choice(kinds)
+ if kind == "recolor":
+ index = rng.randrange(len(edited["groups"]))
+ old = edited["groups"][index]["color"]
+ shape = edited["groups"][index]["shape"]
+ taken = {
+ g["color"]
+ for i, g in enumerate(edited["groups"])
+ if i != index and g["shape"] == shape
+ }
+ edited["groups"][index]["color"] = rng.choice(
+ [c for c in COLORS if c != old and c not in taken]
+ )
+ detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["color"]}
+ elif kind == "count":
+ index = rng.randrange(len(edited["groups"]))
+ old = edited["groups"][index]["count"]
+ edited["groups"][index]["count"] = old % 4 + 1
+ detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["count"]}
+ elif kind == "remove":
+ index = rng.randrange(len(edited["groups"]))
+ edited["groups"].pop(index)
+ edited["relations"] = [
+ r for r in edited["relations"] if r["a"] != index and r["b"] != index
+ ]
+ for relation in edited["relations"]:
+ relation["a"] -= relation["a"] > index
+ relation["b"] -= relation["b"] > index
+ detail = {"kind": kind, "group": index}
+ else:
+ index = rng.randrange(len(edited["relations"]))
+ old = edited["relations"][index]["relation"]
+ flip = {"left of": "right of", "right of": "left of", "above": "below", "below": "above"}
+ edited["relations"][index]["relation"] = flip[old]
+ detail = {"kind": kind, "relation_index": index, "from": old, "to": flip[old]}
+ return edited, detail
+
+
+def main() -> None:
+ args = parse_args()
+ rng = random.Random(args.seed)
+ output = Path(args.output_dir)
+ (output / "images").mkdir(parents=True, exist_ok=True)
+
+ total = args.vision_only + args.text_only + args.val + args.test
+ scenes, captions = [], []
+ while len(scenes) < total:
+ scene = sample_scene(rng, args)
+ if place_groups(scene, rng, args.image_size) is None:
+ continue
+ scenes.append(scene)
+ captions.append(
+ [
+ caption_scene(scene, rng, args.report_bias)
+ for _ in range(args.captions)
+ ]
+ )
+
+ rows = list(range(total))
+ vision_only = rows[: args.vision_only]
+ text_only = rows[args.vision_only : args.vision_only + args.text_only]
+ val = rows[args.vision_only + args.text_only : args.vision_only + args.text_only + args.val]
+ test = rows[-args.test :]
+
+ needs_render = set(vision_only) | set(val) | set(test)
+ rendered = 0
+ for row in sorted(needs_render):
+ for view in range(args.visual_views):
+ image = None
+ while image is None:
+ image = render_scene(
+ scenes[row], rng, args.image_size, texture=args.texture
+ )
+ image.save(output / "images" / f"scene{row:06d}_v{view}.png")
+ rendered += 1
+
+ interventions = []
+ for row in val + test:
+ for _ in range(args.interventions):
+ edited, detail = intervene(scenes[row], rng)
+ if place_groups(edited, rng, args.image_size) is None:
+ continue
+ index = len(interventions)
+ image = None
+ while image is None:
+ image = render_scene(
+ edited, rng, args.image_size, texture=args.texture
+ )
+ image.save(output / "images" / f"edit{index:06d}.png")
+ interventions.append(
+ {
+ "index": index,
+ "row": row,
+ "detail": detail,
+ "caption": caption_scene(edited, rng, args.report_bias),
+ }
+ )
+
+ vocabulary = sorted(
+ {
+ token
+ for caption_list in captions
+ for caption in caption_list
+ for token in caption.replace(",", " ").replace(".", " ").split()
+ }
+ )
+ manifest = {
+ "dataset": "synth_v0",
+ "image_dir": str(output / "images"),
+ "image_size": args.image_size,
+ "seed": args.seed,
+ "visual_views": args.visual_views,
+ "vision_only_train": vision_only,
+ "text_only_train": text_only,
+ "paired_train": [],
+ "val": val,
+ "test": test,
+ "all_rows": total,
+ "vocabulary_size": len(vocabulary),
+ "vocabulary": vocabulary,
+ "protocol": (
+ "vision_only_train and text_only_train are disjoint scene rows; "
+ "val/test pairs are held out for evaluation only. Captions "
+ "mention exactly the discrete scene state."
+ ),
+ }
+ write_json(output / "manifest.json", manifest)
+ write_json(output / "scenes.private.json", {"scenes": scenes})
+ write_json(output / "captions.json", {"captions": captions})
+ write_json(output / "interventions.private.json", {"interventions": interventions})
+ print(
+ json.dumps(
+ {
+ "scenes": total,
+ "rendered_images": rendered,
+ "interventions": len(interventions),
+ "vocabulary_size": len(vocabulary),
+ "example_caption": captions[test[0]][0],
+ }
+ )
+ )
+
+
+if __name__ == "__main__":
+ main()