diff --git a/autoarray/structures/triangles/coordinate_array.py b/autoarray/structures/triangles/coordinate_array.py index eab8e3497..30114e620 100644 --- a/autoarray/structures/triangles/coordinate_array.py +++ b/autoarray/structures/triangles/coordinate_array.py @@ -1,4 +1,6 @@ from abc import ABC +from functools import lru_cache +from typing import Optional, Tuple import numpy as np @@ -7,6 +9,100 @@ from autoarray.structures.triangles.array import ArrayTriangles +def _lattice_coordinates( + y_min: float, y_max: float, x_min: float, x_max: float, scale: float +) -> np.ndarray: + """ + The integer ``(y, x)`` lattice coordinates `CoordinateArrayTriangles.for_limits_and_scale` + tiles the rectangle with, as an ``(N, 2)`` int array. + """ + y_shift = int(2 * y_min / scale) + x_shift = int(x_min / (HEIGHT_FACTOR * scale)) + + coordinates = [] + + for y in range(y_shift, int(2 * y_max / scale) + 1): + for x in range(x_shift - 1, int(x_max / (HEIGHT_FACTOR * scale)) + 2): + coordinates.append([y, x]) + + return np.array(coordinates) + + +@lru_cache(maxsize=32) +def static_vertex_table( + y_min: float, y_max: float, x_min: float, x_max: float, scale: float +) -> Tuple[np.ndarray, np.ndarray]: + """ + The geometrically unique vertices of the initial triangle lattice + `CoordinateArrayTriangles.for_limits_and_scale` builds, plus the index map from each + triangle's three vertices into them. + + The lattice is fixed by its limits and scale, so it is known before any tracing: on the JAX + `PointSolver` path its vertices are a compile-time constant, and only the geometrically + distinct ones need deflecting. Of the ``3N`` vertex slots most are shared by up to six + triangles (for the ``+-9.9"`` / ``0.2"`` lattice: 69 849 slots, 11 859 distinct points). They + cannot be deduplicated on the floats -- the same point computed from two neighbouring + triangle centres can differ by one ulp (28 665 exact-float distinct rows) -- so they are keyed + on the integer lattice position instead. Vertex ``k`` of the triangle at integer coordinates + ``(cy, cx)`` with flip ``f = +-1`` sits at ``(0.5 * s * (cy + f * dy), 0.5 * h * s * (2 * cx + + f * dx))`` with ``(dy, dx)`` in ``((0, 1), (1, -1), (-1, -1))``, ``s`` the side length and ``h`` + the height factor; ``(cy + f * dy, 2 * cx + f * dx)`` is therefore an exact integer key. + + Each unique vertex takes the float value of its first occurrence, computed with the same + arithmetic as `CoordinateArrayTriangles.triangles`, so ``vertices[indices]`` equals + ``triangles`` to within an ulp per element and is bit-identical for every first occurrence. + + Built in NumPy (never staged into a JAX trace) and cached per geometry; the arrays are + read-only because the cache hands the same objects to every caller. + + Parameters + ---------- + y_min, y_max, x_min, x_max + The limits of the rectangle the lattice tiles. + scale + The side length of the triangles. + + Returns + ------- + ``(vertices, indices)``: the ``(V, 2)`` float64 unique vertices and the ``(N, 3)`` int index map + such that ``vertices[indices]`` is the ``(N, 3, 2)`` triangle array. + """ + coordinates = _lattice_coordinates(y_min, y_max, x_min, x_max, scale) + + flip = np.where((coordinates[:, 0] + coordinates[:, 1]) % 2 != 0, -1, 1)[:, None] + + # Same operation order as `CoordinateArrayTriangles.centres` / `.triangles` (zero offsets). + scaling_factors = np.array([0.5 * scale, HEIGHT_FACTOR * scale]) + centres = scaling_factors * coordinates + np.array([0.0, 0.0]) + triangles = np.stack( + ( + centres + flip * np.array([0.0, 0.5 * scale * HEIGHT_FACTOR]), + centres + flip * np.array([0.5 * scale, -0.5 * scale * HEIGHT_FACTOR]), + centres + flip * np.array([-0.5 * scale, -0.5 * scale * HEIGHT_FACTOR]), + ), + axis=1, + ) + + offsets = np.array([[0, 1], [1, -1], [-1, -1]]) + keys = np.stack( + ( + coordinates[:, None, 0] + flip * offsets[None, :, 0], + 2 * coordinates[:, None, 1] + flip * offsets[None, :, 1], + ), + axis=-1, + ).reshape(-1, 2) + + _, first, inverse = np.unique(keys, axis=0, return_index=True, return_inverse=True) + + vertices = np.ascontiguousarray(triangles.reshape(-1, 2)[first], dtype=np.float64) + indices = np.ascontiguousarray(inverse.reshape(-1, 3)) + + vertices.setflags(write=False) + indices.setflags(write=False) + + return vertices, indices + + class CoordinateArrayTriangles(AbstractTriangles, ABC): def __init__( @@ -16,6 +112,7 @@ def __init__( x_offset: float = 0.0, y_offset: float = 0.0, flipped: bool = False, + vertex_table: Optional[Tuple[np.ndarray, np.ndarray]] = None, ): """ Represents a set of triangles by integer coordinates. @@ -30,10 +127,18 @@ def __init__( Whether the triangles are flipped upside down. y_offset An y_offset to apply to the y coordinates so that up-sampled triangles align. + vertex_table + An optional precomputed ``(vertices, indices)`` pair (see `static_vertex_table`) that + `vertices` and `indices` return instead of the flat per-triangle table. It must describe + exactly these ``coordinates``; only `for_limits_and_scale(..., static_vertices=True)` + sets it. Derived lattices (`for_indexes`, `up_sample`, `neighborhood`) do not inherit + it, and it is not part of the pytree (`tree_flatten`), so an unflattened copy falls + back to the flat table -- which is still correct, only unshared. """ import jax.numpy as jnp self.coordinates = coordinates + self.vertex_table = vertex_table self.side_length = side_length self.flipped = flipped @@ -51,6 +156,7 @@ def for_limits_and_scale( x_min: float, x_max: float, scale: float = 1.0, + static_vertices: bool = False, **_, ): """ @@ -75,21 +181,27 @@ def for_limits_and_scale( The limits of the rectangle to tile. scale The side length of the triangles. + static_vertices + If ``True``, attach the cached `static_vertex_table` for this geometry, so `vertices` + is the ``(V, 2)`` table of geometrically unique lattice vertices (11 859 rows rather + than 69 849 for the ``+-9.9"`` / ``0.2"`` lattice) and `indices` maps each triangle + into it. Consumers that deflect `vertices` then evaluate each lattice point once. The + limits and scale must be concrete Python / NumPy numbers (they are the cache key). """ import jax.numpy as jnp - y_shift = int(2 * y_min / scale) - x_shift = int(x_min / (HEIGHT_FACTOR * scale)) - - coordinates = [] - - for y in range(y_shift, int(2 * y_max / scale) + 1): - for x in range(x_shift - 1, int(x_max / (HEIGHT_FACTOR * scale)) + 2): - coordinates.append([y, x]) + vertex_table = None + if static_vertices: + vertex_table = static_vertex_table( + float(y_min), float(y_max), float(x_min), float(x_max), float(scale) + ) return cls( - coordinates=jnp.array(coordinates), + coordinates=jnp.array( + _lattice_coordinates(y_min, y_max, x_min, x_max, scale) + ), side_length=scale, + vertex_table=vertex_table, ) def tree_flatten(self): @@ -294,9 +406,17 @@ def _vertices_and_indices(self): autolens_profiling#297). NaN padding rows (from `for_indexes`) trace to NaN triangles, which every `Shape.mask` rejects, so containment is unchanged. The NumPy sibling `CoordinateArrayTrianglesNp` still deduplicates, because its shapes are dynamic. + + When a `vertex_table` is attached (the static initial lattice, see `static_vertex_table`) + it is returned instead: a compile-time-constant ``(V, 2)`` table of the geometrically + unique vertices and its ``(N, 3)`` index map. """ import jax.numpy as jnp + if self.vertex_table is not None: + vertices, indices = self.vertex_table + return jnp.asarray(vertices), jnp.asarray(indices) + flat = self.triangles.reshape(-1, 2) indices = jnp.arange(flat.shape[0]).reshape(-1, 3) return flat, indices diff --git a/test_autoarray/structures/triangles/test_coordinate_jax.py b/test_autoarray/structures/triangles/test_coordinate_jax.py index 10085f60f..75e44ea5d 100644 --- a/test_autoarray/structures/triangles/test_coordinate_jax.py +++ b/test_autoarray/structures/triangles/test_coordinate_jax.py @@ -7,6 +7,12 @@ pin that the table round-trips to the triangles exactly, that NaN padding stays NaN and is never returned by containment, that the kept triangles match the deduplicating NumPy sibling, and that the traced containment carries no sort. + +The static initial lattice (`for_limits_and_scale(..., static_vertices=True)`, point-source CPU +phase 3) instead carries the cached `static_vertex_table` of geometrically unique vertices: the +tests below pin its size for the PointSolver lattice (11 859 of 69 849 slots), that it gathers back +to the triangles within a few ulp, that containment matches the flat table and the NumPy sibling, +that derived lattices drop it, and that the cache is keyed on geometry and read-only. """ import importlib @@ -16,6 +22,13 @@ import pytest +# Plain data, defined outside the jax guard: the module-level @parametrize decorators read them +# at collection time, including on the NumPy-only matrix env. +LIMITS = dict(y_min=-1.0, y_max=1.0, x_min=-1.0, x_max=1.0, scale=0.5) + +# The PointSolver's default image-plane extent for a 100x100, 0.2" grid. +SOLVER_LIMITS = dict(y_min=-9.9, y_max=9.9, x_min=-9.9, x_max=9.9, scale=0.2) + # jax is an `[optional]` extra and is absent on the NumPy-only matrix env: every test in this module # skips there. if importlib.util.find_spec("jax") is None: @@ -33,17 +46,26 @@ def test__placeholder_requires_jax(): # pragma: no cover from autoarray.structures.triangles.array import MAX_CONTAINING_SIZE from autoarray.structures.triangles.coordinate_array import ( CoordinateArrayTriangles, + static_vertex_table, ) from autoarray.structures.triangles.coordinate_array_np import ( CoordinateArrayTrianglesNp, ) from autoarray.structures.triangles.shape import Point - LIMITS = dict(y_min=-1.0, y_max=1.0, x_min=-1.0, x_max=1.0, scale=0.5) - def _lattice(): return CoordinateArrayTriangles.for_limits_and_scale(**LIMITS) + def _static_lattice(limits=LIMITS): + return CoordinateArrayTriangles.for_limits_and_scale( + **limits, static_vertices=True + ) + + def _max_ulps(a, b): + a = np.asarray(a) + b = np.asarray(b) + return float(np.max(np.abs(a - b) / np.spacing(np.maximum(np.abs(a), 1.0)))) + def _lattice_np(): return CoordinateArrayTrianglesNp.for_limits_and_scale(**LIMITS) @@ -175,3 +197,135 @@ def test_no_sort(): sorts = re.findall(r"\bsort\b", compiled, flags=re.IGNORECASE) assert not sorts, f"traced containment carries {len(sorts)} sort op(s)" + + +def test__default_lattice_has_no_vertex_table(): + assert _lattice().vertex_table is None + + +def test__static_vertex_table__solver_lattice_counts(): + """ + The +-9.9" / 0.2" lattice has 23 283 triangles, 69 849 vertex slots, 28 665 exact-float distinct + rows and 11 859 geometrically distinct points -- the same count the deduplicating NumPy sibling + gives once its ulp-level duplicates are rounded together. + """ + lattice = _static_lattice(SOLVER_LIMITS) + vertices = np.asarray(lattice.vertices) + indices = np.asarray(lattice.indices) + + assert lattice.coordinates.shape[0] == 23283 + assert vertices.shape == (11859, 2) + assert indices.shape == (23283, 3) + assert indices.min() == 0 and indices.max() == 11858 + assert np.unique(np.round(vertices, 9), axis=0).shape[0] == 11859 + + triangles_np = CoordinateArrayTrianglesNp.for_limits_and_scale(**SOLVER_LIMITS) + flat_np = np.asarray(triangles_np.triangles).reshape(-1, 2) + assert np.unique(flat_np, axis=0).shape[0] == 28665 + assert np.unique(np.round(flat_np, 9), axis=0).shape[0] == 11859 + + +@pytest.mark.parametrize("limits", [LIMITS, SOLVER_LIMITS]) +def test__static_vertex_table_round_trips_to_triangles(limits): + static = _static_lattice(limits) + default = CoordinateArrayTriangles.for_limits_and_scale(**limits) + + assert np.array_equal( + np.asarray(static.coordinates), np.asarray(default.coordinates) + ) + + gathered = np.asarray(static.vertices)[np.asarray(static.indices)] + assert gathered.shape == np.asarray(default.triangles).shape + assert _max_ulps(gathered, default.triangles) <= 4 + + via_with_vertices = np.asarray(static.with_vertices(static.vertices).triangles) + assert np.array_equal(via_with_vertices, gathered) + + +@pytest.mark.parametrize("point_index", range(5)) +def test__static_vertices_kept_triangles_match_numpy(point_index): + triangles_np = _lattice_np() + point = _source_points(triangles_np)[point_index] + + kept_np = _kept_rows(triangles_np, point) + kept = _kept_rows(_static_lattice(), point) + + assert 0 < len(kept_np) < MAX_CONTAINING_SIZE + assert kept == kept_np + + +@pytest.mark.parametrize("point_index", range(5)) +def test__static_vertices_containing_indices_match_default(point_index): + point = Point(*_source_points(_lattice_np())[point_index]) + static = _static_lattice() + default = _lattice() + + static_indices = np.asarray( + static.with_vertices(static.vertices).containing_indices(point) + ) + default_indices = np.asarray( + default.with_vertices(default.vertices).containing_indices(point) + ) + + assert set(static_indices[static_indices >= 0]) == set( + default_indices[default_indices >= 0] + ) + + +def test__static_vertices_containing_indices__jit_matches_eager(): + static = _static_lattice(SOLVER_LIMITS) + point = Point(0.13, -0.27) + + def containing(): + return static.with_vertices(static.vertices).containing_indices(point) + + assert np.array_equal(np.asarray(jax.jit(containing)()), np.asarray(containing())) + + +def test__derived_lattices_drop_the_vertex_table(): + static = _static_lattice() + n = static.coordinates.shape[0] + + for derived in ( + static.for_indexes(jnp.arange(4)), + static.up_sample(), + static.neighborhood(), + ): + assert derived.vertex_table is None + assert derived.vertices.shape == (3 * derived.coordinates.shape[0], 2) + + assert static.vertices.shape[0] < 3 * n + + +def test__static_vertex_table_is_cached_per_geometry(): + first = static_vertex_table(-1.0, 1.0, -1.0, 1.0, 0.5) + + assert static_vertex_table(-1.0, 1.0, -1.0, 1.0, 0.5) is first + assert _static_lattice().vertex_table is first + + finer = static_vertex_table(-1.0, 1.0, -1.0, 1.0, 0.25) + assert finer is not first + assert finer[0].shape[0] > first[0].shape[0] + + +def test__static_vertex_table_is_read_only(): + vertices, indices = static_vertex_table(-1.0, 1.0, -1.0, 1.0, 0.5) + + assert not vertices.flags.writeable + assert not indices.flags.writeable + with pytest.raises(ValueError): + vertices[0, 0] = 1.0 + with pytest.raises(ValueError): + indices[0, 0] = 1 + + +def test_no_sort__static_vertices(): + static = _static_lattice(SOLVER_LIMITS) + + def containing(): + return static.with_vertices(static.vertices).containing_indices(Point(0.1, 0.2)) + + compiled = jax.jit(containing).lower().compile().as_text() + + sorts = re.findall(r"\bsort\b", compiled, flags=re.IGNORECASE) + assert not sorts, f"traced containment carries {len(sorts)} sort op(s)"