torch-dimensions / tests /test_fuzz.py
Celsia's picture
Upload folder using huggingface_hub
ecc81b3 verified
Raw
History Blame Contribute Delete
8.31 kB
"""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}"