qwenjev / data /augment.py
tchbcb's picture
QwenJev: multimodal-retrofitted NanoJev for ARC-AGI-3 (initial skeleton)
3e04895 verified
Raw History Blame Contribute Delete
2.53 kB
"""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