"""Grid augmentations for the QwenJev training set. ARC tradition (AGI-1/2) shows permutation robustness transfers: transpose, mirror, rotate, recolour. Applied consistently to state + candidates, they multiply effective data without touching game semantics. ACTION6 coordinates are transformed together with the grid - get this wrong and labels lie. """ from __future__ import annotations import random from typing import Sequence Grid = Sequence[Sequence[int]] def transpose(g: Grid) -> list[list[int]]: return [list(row) for row in zip(*g)] def mirror_x(g: Grid) -> list[list[int]]: return [list(reversed(row)) for row in g] def mirror_y(g: Grid) -> list[list[int]]: return [list(row) for row in reversed(g)] def rot90(g: Grid) -> list[list[int]]: """Clockwise 90 degrees.""" h, w = len(g), len(g[0]) return [[g[h - 1 - y][x] for y in range(h)] for x in range(w)] def recolour(g: Grid, rng: random.Random) -> list[list[int]]: """Permute the 16 cell colours (fixing 0 = background).""" perm = list(range(16)) rest = perm[1:] rng.shuffle(rest) perm[1:] = rest return [[perm[int(v)] for v in row] for row in g] def coord_transform(name: str, x: int, y: int, h: int, w: int) -> tuple[int, int]: """Apply the same geometric op to an ACTION6 (x, y) candidate.""" if name == "mirror_x": return w - 1 - x, y if name == "mirror_y": return x, h - 1 - y if name == "rot90": return h - 1 - y, x # clockwise return x, y def augment_grid(g: Grid, rng: random.Random) -> tuple[str, list[list[int]]]: """Return (op_name, transformed_grid).""" ops = { "identity": lambda g: [list(r) for r in g], "mirror_x": mirror_x, "mirror_y": mirror_y, "rot90": rot90, "transpose": transpose, "recolour": lambda g: recolour(g, rng), } name = rng.choice(list(ops)) return name, ops[name](g) def augment_sample(sample: dict, rng: random.Random) -> dict: """Augment one training sample (choice family only uses geometry here).""" name, g = augment_grid(sample["state_grids"], rng) if "state_grids" in sample \ else ("identity", sample.get("grids")) h, w = len(g), len(g[0]) if g else (0, 0) out = dict(sample) if name != "identity": out["state"] = f"aug:{name}\n{sample['state']}" if "x" in out and out.get("x") is not None: out["x"], out["y"] = coord_transform(name, int(out["x"]), int(out["y"]), h, w) return out