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
132 changes: 130 additions & 2 deletions autoarray/structures/triangles/array.py
Original file line number Diff line number Diff line change
@@ -1,19 +1,69 @@
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__(
self,
indices,
vertices,
max_containing_size=MAX_CONTAINING_SIZE,
step0_layout: Optional[Step0Layout] = None,
**kwargs,
):
"""
Expand All @@ -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)
Expand Down Expand Up @@ -159,14 +215,84 @@ 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,
size=self.max_containing_size,
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
Expand Down Expand Up @@ -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):
Expand All @@ -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):
Expand All @@ -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],
)


Expand Down
129 changes: 120 additions & 9 deletions autoarray/structures/triangles/coordinate_array.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -83,24 +84,118 @@ 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],
),
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):
Expand All @@ -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.
Expand All @@ -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

Expand Down Expand Up @@ -191,17 +293,25 @@ 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(
_lattice_coordinates(y_min, y_max, x_min, x_max, scale)
),
side_length=scale,
vertex_table=vertex_table,
step0_layout=step0_layout,
)

def tree_flatten(self):
Expand Down Expand Up @@ -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":
Expand Down
Loading
Loading