Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
138 changes: 129 additions & 9 deletions autoarray/structures/triangles/coordinate_array.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
from abc import ABC
from functools import lru_cache
from typing import Optional, Tuple

import numpy as np

Expand All @@ -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__(
Expand All @@ -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.
Expand All @@ -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

Expand All @@ -51,6 +156,7 @@ def for_limits_and_scale(
x_min: float,
x_max: float,
scale: float = 1.0,
static_vertices: bool = False,
**_,
):
"""
Expand All @@ -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):
Expand Down Expand Up @@ -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
Expand Down
158 changes: 156 additions & 2 deletions test_autoarray/structures/triangles/test_coordinate_jax.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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)

Expand Down Expand Up @@ -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)"
Loading