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