File size: 8,305 Bytes
ecc81b3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
"""Seeded fuzz over the library's invariants, checked against slow references.

Targeted tests check configurations someone thought of; these check the ones
nobody did. Every case is seeded, so a failure reproduces exactly — paste the
printed config into a targeted test and it stays failed until fixed.
"""

import torch

import torch_dimensions as td
from torch_dimensions.compose.kernel import axial_contract
from torch_dimensions.compose.scan import axial_apply
from torch_dimensions.data.coords import from_coords
from torch_dimensions.data.window import LatticeWindow

_REL = 1e-3  # mirror of the kernel module's cancellation threshold


def _rand_lattice(g, rank, time, sparse):
    shape = tuple(int(torch.randint(1, 5, (1,), generator=g)) for _ in range(rank))
    valid = None
    if sparse and rank > 0:
        valid = torch.rand(shape, generator=g) > 0.5
        if not valid.any():
            valid.reshape(-1)[int(torch.randint(0, valid.numel(), (1,), generator=g))] = True
    return td.Lattice(shape=shape, valid=valid, time=time)


def test_fold_scatter_and_permutation_round_trip_on_random_lattices():
    g = torch.Generator().manual_seed(0)
    for i in range(60):
        rank = int(torch.randint(1, 5, (1,), generator=g))
        time = bool(torch.randint(0, 2, (1,), generator=g))
        sparse = bool(torch.randint(0, 2, (1,), generator=g))
        lat = _rand_lattice(g, rank, time, sparse)
        lead = (2, 3) if time else (2,)
        x = torch.randn(*lead, *lat.shape, 4, generator=g)
        for axis in range(lat.n_axes):
            seq, restore = lat.to_sequence(x, axis)
            assert torch.equal(lat.from_sequence(seq, restore), x), f"[{i}] axis {axis} {lat}"
            perm, inv = lat.permutation(axis)
            assert list(torch.argsort(torch.tensor(perm))) == list(inv), f"[{i}] {lat}"
        xm = x * lat.mask().to(x.dtype)
        assert torch.equal(lat.scatter(lat.gather(xm)), xm), f"[{i}] {lat}"
        assert lat.flat_idx.numel() == lat.n_valid, f"[{i}] {lat}"


def test_axial_contract_matches_a_per_line_loop_on_random_sparse_lattices():
    """Independent reference: an explicit loop over every line, including the
    relative-cancellation rule for degenerate denominators."""
    g = torch.Generator().manual_seed(1)
    for i in range(30):
        rank = int(torch.randint(1, 4, (1,), generator=g))
        lat = _rand_lattice(g, rank, False, True)
        axis = int(torch.randint(0, rank, (1,), generator=g))
        a_len = lat.axis_size(axis)
        mask = lat.mask().to(torch.float64)
        x = torch.randn(2, *lat.shape, 3, dtype=torch.float64, generator=g) * mask
        kernel = torch.randn(a_len, a_len, dtype=torch.float64, generator=g)  # signed
        got = axial_contract(x, lat, axis, kernel, valid=mask)

        seq, restore = lat.to_sequence(x, axis)
        mseq, _ = lat.to_sequence(mask.expand(*x.shape[:-1], 1), axis)
        out = torch.zeros_like(seq)
        for m in range(seq.shape[0]):
            pres = mseq[m, :, 0]
            for q in range(a_len):
                den = float((kernel[q] * pres).sum())
                den_abs = float((kernel[q].abs() * pres).sum())
                num = (kernel[q].unsqueeze(-1) * seq[m] * pres.unsqueeze(-1)).sum(0)
                out[m, q] = num if abs(den) <= _REL * den_abs else num / den
        want = lat.from_sequence(out, restore)
        assert torch.allclose(got, want, atol=1e-10), f"[{i}] {lat} axis {axis}"


def test_axial_apply_matches_cumsum_on_random_configs():
    g = torch.Generator().manual_seed(2)
    for i in range(40):
        rank = int(torch.randint(1, 5, (1,), generator=g))
        time = bool(torch.randint(0, 2, (1,), generator=g))
        lat = _rand_lattice(g, rank, time, False)
        lead = (2, 3) if time else (2,)
        x = torch.randn(*lead, *lat.shape, 3, generator=g)
        axis = int(torch.randint(0, lat.n_axes, (1,), generator=g))
        rev = bool(torch.randint(0, 2, (1,), generator=g))
        chunk = [None, 1, 7][int(torch.randint(0, 3, (1,), generator=g))]
        d = lat.tensor_dim(axis)
        want = x.flip(d).cumsum(dim=d).flip(d) if rev else x.cumsum(dim=d)
        got = axial_apply(x, lat, axis, lambda s: s.cumsum(dim=1), reverse=rev, chunk=chunk)
        assert torch.equal(got, want), f"[{i}] rank {rank} axis {axis} rev {rev} chunk {chunk}"


def test_window_tiling_properties_on_random_configs():
    g = torch.Generator().manual_seed(3)
    for i in range(120):
        n = int(torch.randint(2, 40, (1,), generator=g))
        il = int(torch.randint(1, n + 1, (1,), generator=g))
        hz = int(torch.randint(0, n - il + 1, (1,), generator=g))
        st = int(torch.randint(1, 6, (1,), generator=g))
        w = LatticeWindow(n, input_len=il, horizon=hz, stride=st)
        for win in w:
            assert 0 <= win.x0 < win.x1 <= win.y0 <= win.y1 <= n, f"[{i}] {win} n={n}"
            assert win.x1 - win.x0 == il and win.y1 - win.y0 == hz, f"[{i}] {win}"
        at = int(torch.randint(0, n + 1, (1,), generator=g))
        train, test = w.split(at)
        assert all(win.y1 <= at for win in train), f"[{i}] train crosses cut at {at}"
        assert all(win.x0 >= at for win in test), f"[{i}] test crosses cut at {at}"


def test_coords_encode_decode_round_trip_on_random_tables():
    g = torch.Generator().manual_seed(4)
    for i in range(40):
        k = int(torch.randint(1, 4, (1,), generator=g))
        n_rows = int(torch.randint(1, 30, (1,), generator=g))
        rows = [
            tuple(f"v{int(torch.randint(0, 4, (1,), generator=g))}" for _ in range(k))
            for _ in range(n_rows)
        ]
        cm = from_coords(rows, time=False)
        for row, flat in zip(rows, cm.index.tolist(), strict=True):
            assert cm.decode(flat) == row, f"[{i}] {row} -> {cm.decode(flat)}"
        assert torch.equal(cm.encode(rows), cm.index), f"[{i}] encode != index"


def test_fold_round_trips_at_ranks_five_and_six():
    """The rank-1..4 envelope above found real bugs; the machinery claims to be
    rank-generic, so the envelope should stop where patience does, not where
    the claim does. Sizes stay tiny: rank is the variable under test."""
    g = torch.Generator().manual_seed(5)
    for i in range(20):
        rank = 5 + int(torch.randint(0, 2, (1,), generator=g))
        time = bool(torch.randint(0, 2, (1,), generator=g))
        shape = tuple(int(torch.randint(1, 4, (1,), generator=g)) for _ in range(rank))
        valid = torch.rand(shape, generator=g) > 0.5
        valid.reshape(-1)[0] = True
        lat = td.Lattice(shape=shape, valid=valid, time=time)
        lead = (2, 2) if time else (2,)
        x = torch.randn(*lead, *shape, 2, generator=g)
        for axis in range(lat.n_axes):
            seq, restore = lat.to_sequence(x, axis)
            assert seq.shape[-2] == (x.shape[1] if time and axis == 0 else shape[axis - int(time)])
            assert torch.equal(lat.from_sequence(seq, restore), x), f"[{i}] axis {axis} {lat}"
        xm = x * lat.mask().to(x.dtype)
        assert torch.equal(lat.scatter(lat.gather(xm)), xm), f"[{i}] {lat}"


def test_stress_shapes_that_fuzz_would_have_to_be_lucky_to_draw():
    """Degenerate geometries, as explicit cases rather than as fuzz luck: every
    axis length 1, a single existing cell, and one very long axis."""
    cases = [
        td.Lattice(shape=(1, 1, 1, 1), time=True),
        td.Lattice(shape=(1,), time=False),
        td.Lattice(
            shape=(3, 3),
            valid=torch.eye(3, dtype=torch.bool)[:, [0, 0, 0]]
            & torch.tensor([[True, False, False], [False, False, False], [False, False, False]]),
        ),
        td.Lattice(shape=(10_000, 1), names=("long", "thin")),
        td.Lattice(shape=(1, 7), names=("thin", "wide"), time=True),
    ]
    for lat in cases:
        lead = (1, 2) if lat.time else (1,)
        x = torch.randn(*lead, *lat.shape, 2)
        for axis in range(lat.n_axes):
            seq, restore = lat.to_sequence(x, axis)
            assert torch.equal(lat.from_sequence(seq, restore), x), f"{lat} axis {axis}"
        xm = x * lat.mask().to(x.dtype)
        assert torch.equal(lat.scatter(lat.gather(xm)), xm), f"{lat}"
        for name in lat.names or ():
            assert lat.valid_counts(name).min() >= 1, f"{lat} {name}"