diff --git a/autoarray/structures/triangles/array.py b/autoarray/structures/triangles/array.py index 07dfc4369..0ed7cd4e6 100644 --- a/autoarray/structures/triangles/array.py +++ b/autoarray/structures/triangles/array.py @@ -1,12 +1,61 @@ +from dataclasses import dataclass +from typing import Optional, Tuple + import numpy as np from autoarray.structures.triangles.abstract import HEIGHT_FACTOR from autoarray.structures.triangles.abstract import AbstractTriangles +from autoarray.structures.triangles.shape import Point from autoarray.structures.triangles.shape import Shape +from autoarray.structures.triangles.shape import _barycentric_contains MAX_CONTAINING_SIZE = 15 +# Private A/B switch for the step-0 containment route (point-source CPU phase 4b, +# PyAutoArray#579). It only affects `ArrayTriangles` built by +# `CoordinateArrayTriangles.with_vertices` on the static initial lattice (those carrying a +# `Step0Layout`) and tested against a `Point`; every other containment keeps the general +# `shape.mask(self.triangles)` path. All routes return bit-identical kept indices: +# +# - "gather": the general path -- pad, ``(N, 3, 2)`` gather, no-op NaN ``where``. +# - "nopad": the ``(N, 3, 2)`` gather without the pad / NaN ``where`` (no index is -1). +# - "components": six ``(N,)`` 1-D gathers, one per vertex component; no ``(N, 3, 2)`` array. +# - "structured": strided slices of the traced vertex table reshaped to its lattice rows; no +# gather at all, the boolean mask is interleaved back to triangle order. +# +# The switch is read at trace time: a jitted function must be re-traced (fresh closure plus +# ``jax.clear_caches()``) after changing it. +_STEP0_CONTAINMENT = "structured" + + +@dataclass(frozen=True) +class Step0Layout: + """ + The closed-form layout of the static initial lattice's vertex table (see + `autoarray.structures.triangles.coordinate_array.static_lattice_layout`). + + Every value is a Python int, so a layout is hashable and can sit in pytree aux data: under + ``jit`` / ``vmap`` it is a trace-time constant. + + Attributes + ---------- + n_rows, n_cols + The triangle lattice is ``n_rows x n_cols`` triangles, stored row-major. + grid + ``None`` when the vertex table does not follow the closed-form pattern (only the + "gather" / "nopad" / "components" routes then apply), else + ``(n_pairs, pair_width, pad, corners)``: the ``(V, 2)`` table padded by ``pad`` rows is a + ``(n_pairs, pair_width, 2)`` grid, and ``corners[p][k] = (p0, c0, na, nb)`` says that for + the parity class ``p = 2 * (i % 2) + (j % 2)`` of triangle row ``i`` and column ``j``, + vertex ``k`` of the class's triangle ``(a, b)`` (``i = 2a + i % 2``, ``j = 2b + j % 2``) + is grid entry ``(p0 + a, c0 + b)``. + """ + + n_rows: int + n_cols: int + grid: Optional[Tuple] = None + class ArrayTriangles(AbstractTriangles): def __init__( @@ -14,6 +63,7 @@ def __init__( indices, vertices, max_containing_size=MAX_CONTAINING_SIZE, + step0_layout: Optional[Step0Layout] = None, **kwargs, ): """ @@ -26,10 +76,16 @@ def __init__( with the three indices of the vertices. vertices The vertices of the triangles. + step0_layout + Set only by `CoordinateArrayTriangles.with_vertices` on the static initial lattice: + the static layout of ``indices`` into ``vertices``, which lets `containing_indices` + test a `Point` without materialising the ``(N, 3, 2)`` triangle array (see + `_STEP0_CONTAINMENT`). A plain Python constant (pytree aux data), never traced. """ self._indices = indices self._vertices = vertices self.max_containing_size = max_containing_size + self.step0_layout = step0_layout def __len__(self): return len(self.triangles) @@ -159,7 +215,13 @@ def containing_indices(self, shape: Shape) -> np.ndarray: """ import jax.numpy as jnp - inside = shape.mask(self.triangles) + inside = None + # Exactly `Point`: `Circle`, `Triangle`, `Polygon` and `Square` subclass `Point` but + # override `mask`, so they must keep the general path. + if self.step0_layout is not None and type(shape) is Point: + inside = self._step0_point_mask(shape) + if inside is None: + inside = shape.mask(self.triangles) return jnp.where( inside, @@ -167,6 +229,70 @@ def containing_indices(self, shape: Shape) -> np.ndarray: fill_value=-1, )[0] + def _step0_point_mask(self, point: "Point"): + """ + `Point.mask` of the static initial lattice without the general ``(N, 3, 2)`` gather, by + the route `_STEP0_CONTAINMENT` names; ``None`` selects the general path. + + Every route feeds `_barycentric_contains` the same six component values the general path + takes from ``self.triangles`` (at step 0 no index is -1, so its NaN ``where`` is a no-op), + in the same operation order, so the mask -- and the kept indices -- are bit-identical. + """ + import jax.numpy as jnp + + route = _STEP0_CONTAINMENT + vertices = self.vertices + indices = self.indices + + if route == "nopad": + return point.mask(vertices[indices]) + + if route == "components": + v0 = vertices[:, 0] + v1 = vertices[:, 1] + i0 = indices[:, 0] + i1 = indices[:, 1] + i2 = indices[:, 2] + return _barycentric_contains( + v0[i0], v1[i0], v0[i1], v1[i1], v0[i2], v1[i2], point.x, point.y + ) + + if route == "structured" and self.step0_layout.grid is not None: + n_rows = self.step0_layout.n_rows + n_cols = self.step0_layout.n_cols + n_pairs, pair_width, pad, corners = self.step0_layout.grid + if pad: + vertices = jnp.pad(vertices, ((0, pad), (0, 0))) + grid = vertices.reshape(n_pairs, pair_width, 2) + + half_rows = (n_rows + 1) // 2 + half_cols = (n_cols + 1) // 2 + + classes = [] + for corner in corners: + components = [] + for p0, c0, na, nb in corner: + block = grid[p0 : p0 + na, c0 : c0 + nb] + components += [block[..., 0], block[..., 1]] + na, nb = corner[0][2], corner[0][3] + mask = _barycentric_contains(*components, point.x, point.y) + classes.append( + jnp.pad(mask, ((0, half_rows - na), (0, half_cols - nb))) + ) + + # classes[2 * pi + pj][a, b] is triangle (2a + pi, 2b + pj): interleave to row-major. + interleaved = jnp.stack( + ( + jnp.stack((classes[0], classes[1]), axis=-1), + jnp.stack((classes[2], classes[3]), axis=-1), + ), + axis=1, + ).reshape(2 * half_rows, 2 * half_cols) + + return interleaved[:n_rows, :n_cols].reshape(-1) + + return None + def for_indexes(self, indexes: np.ndarray) -> "ArrayTriangles": """ Create a new ArrayTriangles containing indices and vertices corresponding to the given indexes @@ -311,6 +437,7 @@ def with_vertices(self, vertices: np.ndarray) -> "ArrayTriangles": indices=self.indices, vertices=vertices, max_containing_size=self.max_containing_size, + step0_layout=self.step0_layout, ) def tree_flatten(self): @@ -320,7 +447,7 @@ def tree_flatten(self): return ( self.indices, self.vertices, - ), (self.max_containing_size,) + ), (self.max_containing_size, self.step0_layout) @classmethod def tree_unflatten(cls, aux_data, children): @@ -331,6 +458,7 @@ def tree_unflatten(cls, aux_data, children): indices=children[0], vertices=children[1], max_containing_size=aux_data[0], + step0_layout=aux_data[1], ) diff --git a/autoarray/structures/triangles/coordinate_array.py b/autoarray/structures/triangles/coordinate_array.py index 30114e620..91d736c96 100644 --- a/autoarray/structures/triangles/coordinate_array.py +++ b/autoarray/structures/triangles/coordinate_array.py @@ -7,6 +7,7 @@ from autoarray.structures.triangles.abstract import HEIGHT_FACTOR from autoarray.structures.triangles.abstract import AbstractTriangles from autoarray.structures.triangles.array import ArrayTriangles +from autoarray.structures.triangles.array import Step0Layout def _lattice_coordinates( @@ -83,8 +84,27 @@ def static_vertex_table( axis=1, ) + keys = _vertex_keys(coordinates) + + _, 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 + + +def _vertex_keys(coordinates: np.ndarray) -> np.ndarray: + """ + The exact integer lattice key of every ``(3N)`` vertex slot, ordered as + `static_vertex_table` orders them (see its docstring). + """ + flip = np.where((coordinates[:, 0] + coordinates[:, 1]) % 2 != 0, -1, 1)[:, None] offsets = np.array([[0, 1], [1, -1], [-1, -1]]) - keys = np.stack( + return np.stack( ( coordinates[:, None, 0] + flip * offsets[None, :, 0], 2 * coordinates[:, None, 1] + flip * offsets[None, :, 1], @@ -92,15 +112,90 @@ def static_vertex_table( 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)) +@lru_cache(maxsize=32) +def static_lattice_layout( + y_min: float, y_max: float, x_min: float, x_max: float, scale: float +) -> Step0Layout: + """ + The closed-form layout of `static_vertex_table`'s index map, so the step-0 containment can + read each triangle's vertices by strided slicing instead of a gather (PyAutoArray#579). + + The lattice's triangles are row-major, ``n_rows x n_cols``. The vertex table is sorted on the + integer key ``(ky, kx)``; key rows alternate between two widths ``W0, W1`` (equal when the + lattice has an odd number of columns), so the ``(V, 2)`` table -- padded by one row when the + number of key rows is odd -- is a ``(n_pairs, W0 + W1, 2)`` grid of row pairs. Splitting the + triangles into the four parity classes of their (row, column) index fixes each class's flip + and the parity of every vertex's key row, so vertex ``k`` of the class's triangle ``(a, b)`` + is grid entry ``(p0 + a, c0 + b)``: one contiguous 2-D slice per class and vertex. + + The layout is derived from, and checked element-wise against, the actual index map in NumPy: + if any class did not follow the strided pattern, ``grid`` is ``None`` and the structured + route is not used for this geometry. Built once per geometry (cached) and made only of Python + ints, so it is hashable pytree aux data. + """ + vertices, indices = static_vertex_table(y_min, y_max, x_min, x_max, scale) + coordinates = _lattice_coordinates(y_min, y_max, x_min, x_max, scale) - vertices.setflags(write=False) - indices.setflags(write=False) + ys = np.unique(coordinates[:, 0]) + xs = np.unique(coordinates[:, 1]) + n_rows, n_cols = ys.shape[0], xs.shape[0] - return vertices, indices + grid_coordinates = np.stack(np.meshgrid(ys, xs, indexing="ij"), axis=-1) + if coordinates.shape[0] != n_rows * n_cols or not np.array_equal( + coordinates, grid_coordinates.reshape(-1, 2) + ): + return Step0Layout(n_rows=0, n_cols=0, grid=None) + + unique_keys = np.unique(_vertex_keys(coordinates), axis=0) + _, row_counts = np.unique(unique_keys[:, 0], return_counts=True) + R = row_counts.shape[0] + W0 = int(row_counts[0]) + W1 = int(row_counts[1]) if R > 1 else W0 + if ( + unique_keys.shape[0] != vertices.shape[0] + or np.any(row_counts[0::2] != W0) + or np.any(row_counts[1::2] != W1) + ): + return Step0Layout(n_rows=n_rows, n_cols=n_cols, grid=None) + + # Key rows come in pairs of widths (W0, W1). Viewing the table as (n_pairs, W0 + W1, 2) -- + # padded by one row when R is odd -- puts vertex (row r, column c) at + # (r // 2, (r % 2) * W0 + c). + pair_width = W0 + W1 + n_pairs = (R + 1) // 2 + row_start = np.concatenate(([0], np.cumsum(row_counts)[:-1])) + vertex_key_row = np.repeat(np.arange(R), row_counts) + vertex_key_col = np.arange(vertices.shape[0]) - row_start[vertex_key_row] + vertex_pair = (vertex_key_row // 2)[indices].reshape(n_rows, n_cols, 3) + vertex_offset = ((vertex_key_row % 2) * W0 + vertex_key_col)[indices].reshape( + n_rows, n_cols, 3 + ) + + corners = [] + for pi in (0, 1): + for pj in (0, 1): + corner = [] + for k in range(3): + class_pair = vertex_pair[pi::2, pj::2, k] + class_offset = vertex_offset[pi::2, pj::2, k] + na, nb = class_pair.shape + p0, c0 = int(class_pair[0, 0]), int(class_offset[0, 0]) + expected_pair = p0 + np.arange(na)[:, None] + 0 * class_pair + expected_offset = c0 + np.arange(nb)[None, :] + 0 * class_offset + if not ( + np.array_equal(class_pair, expected_pair) + and np.array_equal(class_offset, expected_offset) + ): + return Step0Layout(n_rows=n_rows, n_cols=n_cols, grid=None) + corner.append((p0, c0, int(na), int(nb))) + corners.append(tuple(corner)) + + return Step0Layout( + n_rows=n_rows, + n_cols=n_cols, + grid=(n_pairs, pair_width, int(n_pairs * pair_width - vertices.shape[0]), tuple(corners)), + ) class CoordinateArrayTriangles(AbstractTriangles, ABC): @@ -113,6 +208,7 @@ def __init__( y_offset: float = 0.0, flipped: bool = False, vertex_table: Optional[Tuple[np.ndarray, np.ndarray]] = None, + step0_layout: Optional[Step0Layout] = None, ): """ Represents a set of triangles by integer coordinates. @@ -134,11 +230,17 @@ def __init__( 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. + step0_layout + The `static_lattice_layout` of ``vertex_table``, set with it by + `for_limits_and_scale(..., static_vertices=True)` and passed by `with_vertices` to the + `ArrayTriangles` it returns, whose `containing_indices` then avoids the ``(N, 3, 2)`` + gather. Dropped exactly where ``vertex_table`` is. """ import jax.numpy as jnp self.coordinates = coordinates self.vertex_table = vertex_table + self.step0_layout = step0_layout if vertex_table is not None else None self.side_length = side_length self.flipped = flipped @@ -191,10 +293,17 @@ def for_limits_and_scale( import jax.numpy as jnp vertex_table = None + step0_layout = None if static_vertices: - vertex_table = static_vertex_table( - float(y_min), float(y_max), float(x_min), float(x_max), float(scale) + geometry = ( + float(y_min), + float(y_max), + float(x_min), + float(x_max), + float(scale), ) + vertex_table = static_vertex_table(*geometry) + step0_layout = static_lattice_layout(*geometry) return cls( coordinates=jnp.array( @@ -202,6 +311,7 @@ def for_limits_and_scale( ), side_length=scale, vertex_table=vertex_table, + step0_layout=step0_layout, ) def tree_flatten(self): @@ -437,6 +547,7 @@ def with_vertices(self, vertices: np.ndarray) -> ArrayTriangles: return ArrayTriangles( indices=self.indices, vertices=vertices, + step0_layout=self.step0_layout, ) def for_indexes(self, indexes: np.ndarray) -> "CoordinateArrayTriangles": diff --git a/test_autoarray/structures/triangles/test_coordinate_jax.py b/test_autoarray/structures/triangles/test_coordinate_jax.py index 75e44ea5d..292e831c7 100644 --- a/test_autoarray/structures/triangles/test_coordinate_jax.py +++ b/test_autoarray/structures/triangles/test_coordinate_jax.py @@ -29,6 +29,25 @@ # 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) +# Step-0 containment routes (`autoarray.structures.triangles.array._STEP0_CONTAINMENT`, +# PyAutoArray#579). "gather" is the general path every other route must reproduce bit-for-bit. +STEP0_ROUTES = ("gather", "nopad", "components", "structured") + +# Fuzz geometries for the step-0 routes, as (name, limits, deflect, n_extra): the solver lattice +# (199 x 117 triangles, odd key-row count -> padded grid), the same lattice with an SIS-like +# deflection of its traced vertex table, and three asymmetric lattices -- alternating key-row +# widths 14/15 (even column count), alternating widths 10/11 with an odd key-row count (padded), +# and an even row count at a non-round scale. ``n_extra`` random points are added to every +# vertex, edge midpoint and centroid of the (possibly deflected) lattice; the solver lattices +# instead sample a subset of those (see `_fuzz_points`). +STEP0_FUZZ_GEOMETRIES = ( + ("solver", SOLVER_LIMITS, False, 2048), + ("solver_deflected", SOLVER_LIMITS, True, 1024), + ("alternating", dict(y_min=-3.1, y_max=2.2, x_min=-1.7, x_max=2.9, scale=0.2), False, 8192), + ("alternating_padded", dict(y_min=-2.1, y_max=1.1, x_min=-0.6, x_max=2.5, scale=0.2), True, 8192), + ("even_rows", dict(y_min=-0.9, y_max=1.7, x_min=-2.2, x_max=0.4, scale=0.13), False, 8192), +) + # 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: @@ -43,6 +62,7 @@ def test__placeholder_requires_jax(): # pragma: no cover jax.config.update("jax_enable_x64", True) + from autoarray.structures.triangles import array as triangles_array from autoarray.structures.triangles.array import MAX_CONTAINING_SIZE from autoarray.structures.triangles.coordinate_array import ( CoordinateArrayTriangles, @@ -329,3 +349,240 @@ def containing(): sorts = re.findall(r"\bsort\b", compiled, flags=re.IGNORECASE) assert not sorts, f"traced containment carries {len(sorts)} sort op(s)" + + + +# --------------------------------------------------------------------------------------------- +# Step-0 containment routes (point-source CPU phase 4b, PyAutoArray#579) +# --------------------------------------------------------------------------------------------- + +# Points per jitted, vmapped call; the last chunk is padded by repeating its final point. +_FUZZ_CHUNK = 1024 + +# Set PYAUTO_STEP0_FUZZ_SCALE (an integer >= 1) to multiply the random and sampled fuzz points, +# e.g. 4 for ~1.5e5 points per route. The default (~5.7e4 points per route, every vertex of the +# solver lattice included) keeps this block to ~30 s on a laptop. +_FUZZ_SCALE = int(__import__("os").environ.get("PYAUTO_STEP0_FUZZ_SCALE", "1")) + + +def _deflected(vertices): + """An SIS-like (theta_E = 1.6) deflection of a vertex table: an irregular traced table.""" + vertices = np.asarray(vertices) + radius = np.sqrt(np.sum(vertices**2, axis=1, keepdims=True) + 1e-2) + return vertices - 1.6 * vertices / radius + + +def _fuzz_setup(limits, deflect): + lattice = _static_lattice(limits) + vertices = np.asarray(lattice.vertices) + if deflect: + vertices = _deflected(vertices) + return lattice, vertices + + +def _fuzz_points(limits, deflect, n_extra): + """ + Every vertex of the (possibly deflected) table -- the exact step-0 ties -- plus edge + midpoints, centroids and uniform random points over the table's bounding box. On the + 23 283-triangle solver lattices the midpoints and centroids are a random subset. + """ + lattice, vertices = _fuzz_setup(limits, deflect) + triangles = vertices[np.asarray(lattice.indices)] + rng = np.random.default_rng(579) + + midpoints = np.concatenate( + [0.5 * (triangles[:, a] + triangles[:, b]) for a, b in ((0, 1), (1, 2), (2, 0))] + ) + centroids = triangles.mean(axis=1) + n_extra = n_extra * _FUZZ_SCALE + if triangles.shape[0] > 5000: + midpoints = midpoints[rng.choice(midpoints.shape[0], n_extra // 2, replace=False)] + centroids = centroids[rng.choice(centroids.shape[0], n_extra // 4, replace=False)] + if deflect: + vertices = vertices[rng.choice(vertices.shape[0], n_extra, replace=False)] + + low, high = triangles.reshape(-1, 2).min(axis=0), triangles.reshape(-1, 2).max(axis=0) + random = rng.uniform(low, high, size=(n_extra, 2)) + + return np.concatenate((vertices, midpoints, centroids, random)) + + +def _containing_by_route(route, lattice, vertices, points, monkeypatch): + """ + `containing_indices` of every point under ``route``, with the vertex table a traced argument + (as on the solver path) under ``jit`` + ``vmap``. A fresh closure and ``jax.clear_caches()`` + force a re-trace, since the route is read at trace time. + """ + monkeypatch.setattr(triangles_array, "_STEP0_CONTAINMENT", route) + jax.clear_caches() + + def containing(vertices, point): + point = Point(point[0], point[1]) + triangles = lattice.with_vertices(vertices) + # The routed mask itself, kept to _FUZZ_WIDE entries: `containing_indices` truncates at + # MAX_CONTAINING_SIZE, which a folded (deflected) table can exceed. + inside = triangles._step0_point_mask(point) + if inside is None: + inside = point.mask(triangles.triangles) + wide = jnp.where(inside, size=_FUZZ_WIDE, fill_value=-1)[0] + return jnp.concatenate( + (triangles.containing_indices(point), wide, jnp.sum(inside)[None]) + ) + + batched = jax.jit(jax.vmap(containing, in_axes=(None, 0))) + vertices = jnp.asarray(vertices) + + n = points.shape[0] + padded = np.concatenate( + (points, np.repeat(points[-1:], (-n) % _FUZZ_CHUNK, axis=0)) + ) + out = [ + np.asarray(batched(vertices, jnp.asarray(padded[i : i + _FUZZ_CHUNK]))) + for i in range(0, padded.shape[0], _FUZZ_CHUNK) + ] + return np.concatenate(out)[:n] + + +_FUZZ_REFERENCE = {} + +# Width of the routed-mask comparison (see `_containing_by_route`). +_FUZZ_WIDE = 64 + + +def _fuzz_reference(name, limits, deflect, n_extra, monkeypatch): + if name not in _FUZZ_REFERENCE: + lattice, vertices = _fuzz_setup(limits, deflect) + points = _fuzz_points(limits, deflect, n_extra) + _FUZZ_REFERENCE[name] = ( + points, + _containing_by_route("gather", lattice, vertices, points, monkeypatch), + ) + return _FUZZ_REFERENCE[name] + + +@pytest.mark.parametrize("name, limits, deflect, n_extra", STEP0_FUZZ_GEOMETRIES) +def test__step0_fuzz_geometries_have_a_structured_layout(name, limits, deflect, n_extra): + """ + Every fuzz geometry takes the structured route for real (no silent fall-back), and the + layout is hashable int-only aux data cached per geometry. + """ + from autoarray.structures.triangles.coordinate_array import static_lattice_layout + + lattice = _static_lattice(limits) + layout = lattice.with_vertices(lattice.vertices).step0_layout + + assert layout is not None and layout.grid is not None + assert layout.n_rows * layout.n_cols == lattice.coordinates.shape[0] + assert hash(layout) == hash(static_lattice_layout(*[float(limits[k]) for k in ( + "y_min", "y_max", "x_min", "x_max", "scale")])) + assert static_lattice_layout( + *[float(limits[k]) for k in ("y_min", "y_max", "x_min", "x_max", "scale")] + ) is layout + + +@pytest.mark.parametrize("route", [r for r in STEP0_ROUTES if r != "gather"]) +@pytest.mark.parametrize("name, limits, deflect, n_extra", STEP0_FUZZ_GEOMETRIES) +def test__step0_route_is_bit_identical_to_gather( + name, limits, deflect, n_extra, route, monkeypatch +): + """ + The bit-identity fuzz: each route keeps exactly the triangles the general gather path keeps, + in the same order, for every vertex (exact ties), edge midpoint, centroid and random point + of each geometry, with the vertex table traced under ``jit`` + ``vmap``. + """ + points, reference = _fuzz_reference(name, limits, deflect, n_extra, monkeypatch) + lattice, vertices = _fuzz_setup(limits, deflect) + + result = _containing_by_route(route, lattice, vertices, points, monkeypatch) + + mismatched = np.flatnonzero(np.any(result != reference, axis=1)) + assert mismatched.size == 0, ( + f"{route} on {name}: {mismatched.size} of {points.shape[0]} points differ from gather, " + f"first at {points[mismatched[0]]}: {result[mismatched[0]]} vs {reference[mismatched[0]]}" + ) + + # The fuzz really exercises ties (vertices kept by several triangles) and misses, and the + # wide comparison covers every contained triangle. + contained = reference[:, -1] + assert np.any(contained >= 2) + assert np.any(contained == 0) + assert np.all(contained < _FUZZ_WIDE) + + +@pytest.mark.parametrize("route", STEP0_ROUTES) +def test__step0_refinement_path_is_unchanged(route, monkeypatch): + """ + Only the static step-0 lattice carries a layout. Every derived lattice of a solver-style + refinement (kept -> neighbourhood -> up-sample) drops it, so its containment takes the + general path on every route and matches the gather path exactly; a non-`Point` shape on the + static lattice takes the general path too. + """ + from autoarray.structures.triangles.shape import Circle + + lattice = _static_lattice(SOLVER_LIMITS) + vertices = _deflected(lattice.vertices) + source = tuple(vertices[4321]) + + def refine(): + step0 = lattice.with_vertices(jnp.asarray(vertices)) + kept = lattice.for_indexes(step0.containing_indices(Point(*source))) + up_sampled = kept.neighborhood().up_sample() + refined = up_sampled.with_vertices(jnp.asarray(_deflected(up_sampled.vertices))) + circle = step0.containing_indices(Circle(source[0], source[1], radius=0.3)) + return ( + [kept, kept.neighborhood(), up_sampled], + refined, + ( + np.asarray(step0.containing_indices(Point(*source))), + np.asarray(refined.containing_indices(Point(*source))), + np.asarray(circle), + ), + ) + + monkeypatch.setattr(triangles_array, "_STEP0_CONTAINMENT", route) + derived, refined, outputs = refine() + + for triangles in derived: + assert triangles.step0_layout is None + assert triangles.vertex_table is None + assert refined.step0_layout is None + + monkeypatch.setattr(triangles_array, "_STEP0_CONTAINMENT", "gather") + _, _, expected = refine() + + for output, reference in zip(outputs, expected): + assert np.array_equal(output, reference) + assert np.sum(outputs[0] >= 0) >= 2 + + +def test__step0_default_route_has_no_triangle_gather(): + """ + The HLO guard: the optimised HLO of the default step-0 containment on the solver lattice + carries no gather that materialises the ``(23 283, 3, 2)`` triangle array (XLA emits it as a + ``f64[69849,1,2]`` gather on the general path). The vertex table is a traced argument -- a + closure would let XLA constant-fold the gather away. Red on main (the general path). + """ + lattice = _static_lattice(SOLVER_LIMITS) + n = lattice.coordinates.shape[0] + + def containing(vertices, point): + return lattice.with_vertices(vertices).containing_indices( + Point(point[0], point[1]) + ) + + jax.clear_caches() + compiled = ( + jax.jit(containing) + .lower(lattice.vertices, jnp.array([0.1, 0.2])) + .compile() + .as_text() + ) + + gathered = [] + for shape in re.findall(r"= \w+\[([\d,]*)\]\S* gather\(", compiled): + gathered.append(int(np.prod([int(d) for d in shape.split(",") if d]))) + + assert f"f64[{n},3,2]" not in compiled + assert all(size < 3 * n * 2 for size in gathered), ( + f"step-0 containment gathers {gathered} elements; the triangle array is {3 * n * 2}" + )