| """Phase 1 acceptance for Lattice. See PLAN.md. |
| |
| These tests are deliberately heavier than the module's line count justifies: |
| an axis-order bug here is invisible at this layer and presents as a bad model |
| three phases later. |
| """ |
|
|
| import math |
|
|
| import pytest |
| import torch |
| from hypothesis import HealthCheck, given, settings |
| from hypothesis import strategies as st |
|
|
| from torch_dimensions import Lattice |
|
|
| RANKS = [1, 2, 3, 4, 5] |
|
|
|
|
| def _shape(rank: int) -> tuple[int, ...]: |
| return tuple(range(2, 2 + rank)) |
|
|
|
|
| def _tensor(lat: Lattice, batch: int = 2, t: int = 3, h: int = 4) -> torch.Tensor: |
| lead = (batch, t) if lat.time else (batch,) |
| return torch.randn(*lead, *lat.shape, h) |
|
|
|
|
| |
|
|
|
|
| def test_defaults_and_names(): |
| lat = Lattice(shape=(4, 5)) |
| assert lat.rank == 2 and lat.n_axes == 2 |
| assert lat.axis_names == ("dim0", "dim1") |
| assert lat.tensor_ndim == 4 |
| assert lat.is_dense and lat.n_valid == lat.n_cells == 20 |
|
|
|
|
| def test_time_is_a_normal_axis_but_has_no_static_size(): |
| lat = Lattice(shape=(4, 5), names=("h", "w"), time=True) |
| assert lat.n_axes == 3 |
| assert lat.axis_names == ("time", "h", "w") |
| assert lat.tensor_ndim == 5 |
| assert lat.axis_index("time") == 0 |
| assert lat.tensor_dim("h") == 2 |
| assert lat.axis_size("w") == 5 |
| with pytest.raises(ValueError, match="not a lattice axis"): |
| lat.axis_size("time") |
|
|
|
|
| @pytest.mark.parametrize( |
| ("kwargs", "match"), |
| [ |
| ({"shape": ()}, "at least one axis"), |
| ({"shape": (0, 3)}, "must be positive"), |
| ({"shape": (2, 3), "names": ("a",)}, "names for"), |
| ({"shape": (2, 3), "names": ("a", "a")}, "unique"), |
| ({"shape": (2, 3), "names": ("time", "b"), "time": True}, "reserved"), |
| ], |
| ) |
| def test_construction_errors(kwargs, match): |
| with pytest.raises(ValueError, match=match): |
| Lattice(**kwargs) |
|
|
|
|
| def test_valid_mask_errors(): |
| with pytest.raises(ValueError, match="expected"): |
| Lattice(shape=(2, 3), valid=torch.ones(3, 2, dtype=torch.bool)) |
| with pytest.raises(ValueError, match="no cells"): |
| Lattice(shape=(2, 3), valid=torch.zeros(2, 3, dtype=torch.bool)) |
|
|
|
|
| def test_unknown_axis_names_itself(): |
| lat = Lattice(shape=(4, 5), names=("h", "w")) |
| with pytest.raises(KeyError, match="nope"): |
| lat.axis_index("nope") |
| with pytest.raises(IndexError): |
| lat.axis_index(2) |
| assert lat.axis_index(-1) == 1 |
|
|
|
|
| |
|
|
|
|
| @pytest.mark.parametrize("rank", RANKS) |
| @pytest.mark.parametrize("time", [False, True]) |
| def test_inverse_permutation_matches_argsort(rank, time): |
| """The inverse is hand-built in the module; argsort is the independent check.""" |
| lat = Lattice(shape=_shape(rank), time=time) |
| for axis in range(lat.n_axes): |
| perm, inv = lat.permutation(axis) |
| expected = tuple(torch.argsort(torch.tensor(perm)).tolist()) |
| assert inv == expected, f"axis={axis}, perm={perm}" |
|
|
|
|
| @pytest.mark.parametrize("rank", RANKS) |
| @pytest.mark.parametrize("time", [False, True]) |
| def test_sequence_round_trip_is_the_identity(rank, time): |
| lat = Lattice(shape=_shape(rank), time=time) |
| x = _tensor(lat) |
| for axis in range(lat.n_axes): |
| seq, restore = lat.to_sequence(x, axis) |
| assert torch.equal(lat.from_sequence(seq, restore), x), f"axis={axis}" |
|
|
|
|
| @pytest.mark.parametrize("rank", RANKS) |
| def test_folded_shape_is_batch_times_every_other_axis(rank): |
| lat = Lattice(shape=_shape(rank), time=True) |
| b, t, h = 2, 3, 4 |
| x = _tensor(lat, b, t, h) |
| sizes = (t, *lat.shape) |
| for axis in range(lat.n_axes): |
| seq, _ = lat.to_sequence(x, axis) |
| a = sizes[axis] |
| assert seq.shape == (b * math.prod(sizes) // a, a, h) |
|
|
|
|
| def test_sequence_preserves_values_along_the_swept_axis(): |
| """Round-tripping can hide a transposition that cancels itself. Check the |
| swept axis actually carries the data it should.""" |
| lat = Lattice(shape=(3, 4), names=("h", "w")) |
| x = torch.arange(2 * 3 * 4 * 5, dtype=torch.float32).reshape(2, 3, 4, 5) |
| seq, _ = lat.to_sequence(x, "h") |
| assert seq.shape == (2 * 4, 3, 5) |
| |
| assert torch.equal(seq[0], x[0, :, 0, :]) |
|
|
|
|
| def test_shape_validation(): |
| lat = Lattice(shape=(3, 4)) |
| with pytest.raises(ValueError, match="expected a 4-D tensor"): |
| lat.to_sequence(torch.randn(2, 3, 4), 0) |
| with pytest.raises(ValueError, match="lattice dims"): |
| lat.to_sequence(torch.randn(2, 3, 9, 5), 0) |
|
|
|
|
| |
|
|
|
|
| @pytest.mark.parametrize("rank", RANKS) |
| def test_scatter_gather_round_trip_dense(rank): |
| lat = Lattice(shape=_shape(rank)) |
| x = torch.randn(2, lat.n_valid, 4) |
| assert torch.equal(lat.gather(lat.scatter(x)), x) |
|
|
|
|
| @pytest.mark.parametrize("rank", [1, 2, 3, 4]) |
| @pytest.mark.parametrize("time", [False, True]) |
| def test_scatter_gather_round_trip_sparse(rank, time): |
| shape = _shape(rank) |
| torch.manual_seed(rank) |
| valid = torch.rand(shape) > 0.4 |
| valid.reshape(-1)[0] = True |
| lat = Lattice(shape=shape, valid=valid, time=time) |
| lead = (2, 3) if time else (2,) |
| x = torch.randn(*lead, lat.n_valid, 4) |
| dense = lat.scatter(x) |
| assert dense.shape == (*lead, *shape, 4) |
| assert torch.equal(lat.gather(dense), x) |
|
|
|
|
| def test_scatter_zeroes_cells_that_do_not_exist(): |
| valid = torch.tensor([[True, False], [False, True]]) |
| lat = Lattice(shape=(2, 2), valid=valid) |
| dense = lat.scatter(torch.ones(1, 2, 3)) |
| assert torch.equal(dense[0, 0, 1], torch.zeros(3)) |
| assert torch.equal(dense[0, 1, 0], torch.zeros(3)) |
| assert torch.equal(dense[0, 0, 0], torch.ones(3)) |
|
|
|
|
| def test_scatter_rejects_wrong_cell_count(): |
| lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, False], [False, True]])) |
| with pytest.raises(ValueError, match="expected 2 cells"): |
| lat.scatter(torch.ones(1, 4, 3)) |
|
|
|
|
| def test_mask_broadcasts_over_batch_and_features(): |
| valid = torch.tensor([[True, False], [False, True]]) |
| lat = Lattice(shape=(2, 2), valid=valid, time=True) |
| m = lat.mask(torch.float32) |
| assert m.shape == (1, 1, 2, 2, 1) |
| x = torch.randn(2, 3, 2, 2, 4) |
| assert (x * m)[:, :, 0, 1, :].abs().sum() == 0 |
|
|
|
|
| def test_dense_mask_is_all_true(): |
| lat = Lattice(shape=(2, 3)) |
| assert lat.mask().all() and lat.mask().shape == (1, 2, 3, 1) |
|
|
|
|
| @pytest.mark.parametrize("rank", [1, 2, 3]) |
| def test_valid_counts_matches_manual_sum(rank): |
| shape = _shape(rank) |
| torch.manual_seed(7) |
| valid = torch.rand(shape) > 0.3 |
| valid.reshape(-1)[0] = True |
| lat = Lattice(shape=shape, valid=valid) |
| for i in range(rank): |
| others = tuple(j for j in range(rank) if j != i) |
| expected = valid.sum(dim=others) if others else valid.to(torch.long) |
| assert torch.equal(lat.valid_counts(i), expected.float().clamp_min(1.0)) |
|
|
|
|
| def test_valid_counts_dense_is_the_product_of_other_axes(): |
| lat = Lattice(shape=(2, 3, 4)) |
| assert torch.equal(lat.valid_counts(0), torch.full((2,), 12.0)) |
| assert torch.equal(lat.valid_counts(2), torch.full((4,), 6.0)) |
|
|
|
|
| def test_valid_counts_never_zero_so_division_is_safe(): |
| valid = torch.tensor([[True, True], [False, False]]) |
| lat = Lattice(shape=(2, 2), valid=valid) |
| assert torch.equal(lat.valid_counts(0), torch.tensor([2.0, 1.0])) |
|
|
|
|
| |
|
|
|
|
| def test_to_is_a_noop_for_dense_and_copies_for_sparse(): |
| dense = Lattice(shape=(2, 3)) |
| assert dense.to("cpu") is dense |
| sparse = Lattice(shape=(2, 2), valid=torch.tensor([[True, False], [True, True]])) |
| moved = sparse.to("cpu") |
| assert moved is not sparse and moved.n_valid == 3 |
|
|
|
|
| def test_repr_reports_sparsity(): |
| lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, False], [True, True]]), time=True) |
| r = repr(lat) |
| assert "3/4" in r and "time=True" in r |
|
|
|
|
| |
|
|
|
|
| @settings(deadline=None, max_examples=40, suppress_health_check=[HealthCheck.too_slow]) |
| @given( |
| sizes=st.lists(st.integers(min_value=1, max_value=4), min_size=1, max_size=5), |
| time=st.booleans(), |
| density=st.floats(min_value=0.2, max_value=1.0), |
| seed=st.integers(min_value=0, max_value=2**16), |
| ) |
| def test_round_trips_hold_for_arbitrary_lattices(sizes, time, density, seed): |
| shape = tuple(sizes) |
| torch.manual_seed(seed) |
| valid = torch.rand(shape) < density |
| valid.reshape(-1)[0] = True |
| lat = Lattice(shape=shape, valid=valid, time=time) |
|
|
| lead = (2, 3) if time else (2,) |
| sparse = torch.randn(*lead, lat.n_valid, 2) |
| dense = lat.scatter(sparse) |
| assert torch.equal(lat.gather(dense), sparse) |
|
|
| for axis in range(lat.n_axes): |
| perm, inv = lat.permutation(axis) |
| assert inv == tuple(torch.argsort(torch.tensor(perm)).tolist()) |
| seq, restore = lat.to_sequence(dense, axis) |
| assert torch.equal(lat.from_sequence(seq, restore), dense) |
|
|
|
|
| def test_a_lattice_cannot_be_mutated_after_construction(): |
| """Blocks derive buffers from a lattice at construction and `flat_idx` is |
| cached, so a field changed afterwards leaves both stale — `n_valid` would |
| disagree with `flat_idx` and scatter would misplace data.""" |
| lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, True], [True, False]])) |
| assert lat.n_valid == 3 and lat.flat_idx.tolist() == [0, 1, 2] |
| with pytest.raises(AttributeError, match="immutable"): |
| lat.valid = torch.zeros(2, 2, dtype=torch.bool) |
| with pytest.raises(AttributeError, match="immutable"): |
| lat.shape = (9, 9) |
| with pytest.raises(AttributeError, match="immutable"): |
| del lat.valid |
| |
| assert lat.n_valid == 3 and lat.flat_idx.tolist() == [0, 1, 2] |
|
|
|
|
| def test_to_returns_a_new_lattice_rather_than_mutating(): |
| lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, False], [True, True]])) |
| moved = lat.to("cpu") |
| assert moved is not lat |
| assert moved.n_valid == lat.n_valid |
|
|
|
|
| def test_mask_is_a_copy_not_a_view_of_valid(): |
| """mask(bool) used to be a reshaped view of `valid`: writing into "your" |
| mask corrupted the lattice through it.""" |
| lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, True], [True, False]])) |
| lat.mask()[0, 0, 0, 0] = False |
| assert bool(lat.valid[0, 0]), "mutating the returned mask reached lat.valid" |
|
|
|
|
| def test_the_callers_valid_tensor_is_not_aliased(): |
| """A caller who reuses the tensor they passed as `valid` must not desync |
| the lattice's caches — flat_idx would misplace scatter/gather data.""" |
| v = torch.tensor([[True, True], [True, False]]) |
| lat = Lattice(shape=(2, 2), valid=v) |
| assert lat.flat_idx.tolist() == [0, 1, 2] |
| v[0, 0] = False |
| assert lat.n_valid == 3 |
| assert lat.flat_idx.tolist() == [0, 1, 2] |
|
|
|
|
| |
|
|
|
|
| def test_sliced_narrows_axes_and_the_selection_matches_the_sub_lattice(): |
| lat = Lattice(shape=(4, 5), names=("row", "col"), time=True) |
| sub = lat.sliced(row=slice(1, 3), col=[0, 2, 4]) |
| assert sub.lattice.shape == (2, 3) |
| assert sub.lattice.names == ("row", "col") and sub.lattice.time |
| x = torch.arange(2 * 3 * 4 * 5 * 2).reshape(2, 3, 4, 5, 2).float() |
| got = sub.take(x) |
| assert tuple(got.shape) == (2, 3, 2, 3, 2) |
| assert torch.equal(got, x[:, :, 1:3][:, :, :, [0, 2, 4]]) |
|
|
|
|
| def test_sliced_carries_the_validity_mask_through(): |
| valid = torch.tensor([[True, False, True], [False, False, True], [True, True, False]]) |
| lat = Lattice(shape=(3, 3), names=("a", "b"), valid=valid) |
| sub = lat.sliced(a=slice(0, 2)) |
| assert torch.equal(sub.lattice.valid, valid[:2]) |
| assert sub.lattice.n_valid == 3 |
|
|
|
|
| def test_an_unmentioned_axis_is_kept_whole(): |
| lat = Lattice(shape=(3, 4), names=("a", "b")) |
| sub = lat.sliced(b=slice(0, 2)) |
| assert sub.lattice.shape == (3, 2) |
| x = torch.randn(2, 3, 4, 1) |
| assert torch.equal(sub.take(x), x[:, :, :2]) |
|
|
|
|
| def test_slicing_refuses_what_would_change_the_rank_or_empty_an_axis(): |
| lat = Lattice(shape=(3, 4), names=("a", "b"), time=True) |
| with pytest.raises(TypeError, match="slice\\(1, 2\\)"): |
| lat.sliced(a=1) |
| with pytest.raises(ValueError, match="empty"): |
| lat.sliced(a=slice(2, 2)) |
| with pytest.raises(KeyError, match="nope"): |
| lat.sliced(nope=slice(0, 1)) |
| with pytest.raises(IndexError, match="out of range"): |
| lat.sliced(a=[0, 9]) |
| with pytest.raises(ValueError, match="time"): |
| lat.sliced(time=slice(0, 1)) |
|
|
|
|
| def test_slicing_away_every_existing_cell_is_refused(): |
| valid = torch.tensor([[True, True], [False, False]]) |
| lat = Lattice(shape=(2, 2), names=("a", "b"), valid=valid) |
| with pytest.raises(ValueError, match="no existing cells"): |
| lat.sliced(a=slice(1, 2)) |
|
|
|
|
| def test_merge_is_the_inverse_of_sliced(): |
| valid = torch.rand(6, 4) > 0.4 |
| valid[0, 0] = True |
| lat = Lattice(shape=(6, 4), names=("a", "b"), valid=valid, time=True) |
| left = lat.sliced(a=slice(0, 4)) |
| right = lat.sliced(a=slice(4, 6)) |
| back = Lattice.merge([left.lattice, right.lattice], "a") |
| assert back.shape == lat.shape |
| assert torch.equal(back.valid, lat.valid) |
| x = torch.randn(2, 3, 6, 4, 5) |
| rebuilt = torch.cat([left.take(x), right.take(x)], dim=lat.tensor_dim("a")) |
| assert torch.equal(rebuilt, x) |
|
|
|
|
| def test_merging_a_dense_lattice_with_a_sparse_one_keeps_both_claims(): |
| dense = Lattice(shape=(2, 2), names=("a", "b")) |
| sparse = Lattice(shape=(1, 2), names=("a", "b"), valid=torch.tensor([[True, False]])) |
| merged = Lattice.merge([dense, sparse], "a") |
| assert merged.shape == (3, 2) |
| assert merged.valid.tolist() == [[True, True], [True, True], [True, False]] |
|
|
|
|
| def test_merge_refuses_lattices_that_disagree(): |
| a = Lattice(shape=(2, 3), names=("a", "b")) |
| with pytest.raises(ValueError, match="only along"): |
| Lattice.merge([a, Lattice(shape=(2, 4), names=("a", "b"))], "a") |
| with pytest.raises(ValueError, match="names and time"): |
| Lattice.merge([a, Lattice(shape=(2, 3), names=("a", "b"), time=True)], "a") |
| with pytest.raises(ValueError, match="at least one"): |
| Lattice.merge([], "a") |
|
|