File size: 2,532 Bytes
3e04895
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""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