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
117 changes: 94 additions & 23 deletions autoarray/inversion/mesh/interpolator/delaunay.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,15 @@


def scipy_delaunay(points_np, query_points_np, areas_factor):
"""Compute Delaunay simplices (simplices_padded) and Voronoi areas in one call."""
"""Compute the Delaunay simplices (``simplices_padded``), the query-point
mappings, the split-cross points and the barycentric dual areas in one call.

The dual areas are returned as the sixth element: they weight the split
points here, and they are also the exact quadrature weights of the
barycentric interpolant, which
:meth:`~autoarray.inversion.mesh.mesh_geometry.delaunay.MeshGeometryDelaunay.areas_for_magnification`
consumes (PyAutoArray#524).
"""
from scipy.spatial import Delaunay

max_simplices = 2 * points_np.shape[0]
Expand Down Expand Up @@ -60,7 +68,7 @@ def scipy_delaunay(points_np, query_points_np, areas_factor):
delaunay_points=points_np,
)

return points, simplices_padded, mappings, split_points, splitted_mappings
return points, simplices_padded, mappings, split_points, splitted_mappings, areas


# Query points are located in chunks of this size so the (chunk, N) distance
Expand Down Expand Up @@ -287,6 +295,30 @@ def locate_chunk(q_chunk):
return mappings


def _dual_areas_padded_jnp(points, simplices_padded):
"""In-graph barycentric dual areas from the -1 padded simplex table.

A masked scatter-add of ``triangle_area / 3`` into each of a triangle's
three vertices, skipping the padded (-1) rows. Equivalent to
``barycentric_dual_area_from`` on the unpadded simplices, but written with
``.at[].add`` so it stays inside the JIT program and differentiable with
respect to ``points``.
"""
import jax.numpy as jnp

valid = simplices_padded[:, 0] >= 0
s = simplices_padded.clip(min=0)
p0, p1, p2 = points[s[:, 0]], points[s[:, 1]], points[s[:, 2]]
tri_cross = (p1[:, 0] - p0[:, 0]) * (p2[:, 1] - p0[:, 1]) - (
p1[:, 1] - p0[:, 1]
) * (p2[:, 0] - p0[:, 0])
contrib = jnp.where(valid, 0.5 * jnp.abs(tri_cross) / 3.0, 0.0)
areas = jnp.zeros(points.shape[0], dtype=points.dtype)
for k in range(3):
areas = areas.at[s[:, k]].add(contrib)
return areas


def jax_delaunay(points, query_points, areas_factor=0.5):
"""JAX-path Delaunay construction. Only the qhull triangulation runs on
the host (via ``pure_callback``); point location, dual areas and split
Expand All @@ -307,16 +339,7 @@ def jax_delaunay(points, query_points, areas_factor=0.5):
)

# dual areas via masked scatter-add over the padded simplices
valid = simplices_padded[:, 0] >= 0
s = simplices_padded.clip(min=0)
p0, p1, p2 = points[s[:, 0]], points[s[:, 1]], points[s[:, 2]]
tri_cross = (p1[:, 0] - p0[:, 0]) * (p2[:, 1] - p0[:, 1]) - (
p1[:, 1] - p0[:, 1]
) * (p2[:, 0] - p0[:, 0])
contrib = jnp.where(valid, 0.5 * jnp.abs(tri_cross) / 3.0, 0.0)
areas = jnp.zeros(points.shape[0], dtype=points.dtype)
for k in range(3):
areas = areas.at[s[:, k]].add(contrib)
areas = _dual_areas_padded_jnp(points, simplices_padded)

split_points = split_points_from(
points=points,
Expand All @@ -336,7 +359,7 @@ def jax_delaunay(points, query_points, areas_factor=0.5):
xp=jnp,
)

return points, simplices_padded, mappings, split_points, splitted_mappings
return points, simplices_padded, mappings, split_points, splitted_mappings, areas


def barycentric_dual_area_from(
Expand Down Expand Up @@ -441,12 +464,16 @@ def scipy_delaunay_matern(points_np, query_points_np):
"""
Minimal SciPy Delaunay callback for Matérn regularization.

Returns only what’s needed for mapping:
Returns only what’s needed for mapping, plus the dual areas:
- points (tri.points)
- simplices_padded
- mappings: integer array of pixel indices for each query point,
typically of shape (Q, 3), where each row gives the indices of the
Delaunay mesh vertices ("pixels") associated with that query point.
- areas: the barycentric dual area of every vertex. Matérn
regularization does not need them (there are no split points), but the
magnification quadrature does, so they are returned here too
(PyAutoArray#524).
"""
from scipy.spatial import Delaunay

Expand All @@ -472,13 +499,19 @@ def scipy_delaunay_matern(points_np, query_points_np):
delaunay_points=points_np,
)

return points, simplices_padded, mappings
areas = barycentric_dual_area_from(
points,
simplices,
xp=np,
)

return points, simplices_padded, mappings, areas


def jax_delaunay_matern(points, query_points):
"""JAX-path Matérn variant: qhull-only callback + JAX point location,
returning the same minimal (points, simplices_padded, mappings) contract
as ``scipy_delaunay_matern``."""
returning the same minimal (points, simplices_padded, mappings, areas)
contract as ``scipy_delaunay_matern``."""
import jax.numpy as jnp

simplices_padded, simplex_neighbors, vertex_simplex = _jax_delaunay_tables(points)
Expand All @@ -492,7 +525,9 @@ def jax_delaunay_matern(points, query_points):
xp=jnp,
)

return points, simplices_padded, mappings
areas = _dual_areas_padded_jnp(points, simplices_padded)

return points, simplices_padded, mappings, areas


def triangle_area_xp(c0, c1, c2, xp):
Expand Down Expand Up @@ -589,14 +624,34 @@ def pixel_weights_delaunay_from(
class DelaunayInterface:

def __init__(
self, points, simplices, mappings, split_points, splitted_mappings, xp=np
self,
points,
simplices,
mappings,
split_points,
splitted_mappings,
xp=np,
dual_areas=None,
):
"""
Parameters
----------
dual_areas
The barycentric dual area of every mesh vertex (the sum of
``triangle_area / 3`` over the triangles touching it), computed by
whichever Delaunay construction built this interface -- in-graph on
the JAX path, so it is safe to consume inside a ``jax.jit``.
``None`` for interfaces built without them (e.g. the
natural-neighbour interface in ``sibson.py``), in which case the
mesh geometry falls back to a host-side computation.
"""

self.points = points
self.simplices = simplices
self.mappings = mappings
self.split_points = split_points
self.splitted_mappings = splitted_mappings
self.dual_areas = dual_areas

self.xp = xp

Expand Down Expand Up @@ -653,6 +708,20 @@ def __init__(
xp=xp,
)

@property
def dual_areas(self):
"""
The barycentric dual area of every mesh vertex, taken from the Delaunay
construction that already ran for the mappings -- so it is free here,
and on the JAX path it is an in-graph value safe to use inside a
``jax.jit``.

Subclasses that build no Delaunay tables (the kNN interpolators)
override this with ``None``, which makes the mesh geometry compute the
dual areas host-side only if something actually asks for them.
"""
return self.delaunay.dual_areas

@cached_property
def mesh_geometry(self):

Expand All @@ -664,6 +733,7 @@ def mesh_geometry(self):
mesh=self.mesh,
mesh_grid=self.mesh_grid,
data_grid=self.data_grid,
dual_areas=self.dual_areas,
xp=self._xp,
)

Expand Down Expand Up @@ -700,7 +770,7 @@ def delaunay(self) -> "scipy.spatial.Delaunay":

import jax.numpy as jnp

points, simplices, mappings, split_points, splitted_mappings = (
points, simplices, mappings, split_points, splitted_mappings, areas = (
jax_delaunay(
points=self.mesh_grid_xy,
query_points=self.data_grid.over_sampled.array,
Expand All @@ -710,7 +780,7 @@ def delaunay(self) -> "scipy.spatial.Delaunay":

else:

points, simplices, mappings, split_points, splitted_mappings = (
points, simplices, mappings, split_points, splitted_mappings, areas = (
scipy_delaunay(
points_np=self.mesh_grid_xy,
query_points_np=self.data_grid.over_sampled.array,
Expand All @@ -724,14 +794,14 @@ def delaunay(self) -> "scipy.spatial.Delaunay":

import jax.numpy as jnp

points, simplices, mappings = jax_delaunay_matern(
points, simplices, mappings, areas = jax_delaunay_matern(
points=self.mesh_grid_xy,
query_points=self.data_grid.over_sampled.array,
)

else:

points, simplices, mappings = scipy_delaunay_matern(
points, simplices, mappings, areas = scipy_delaunay_matern(
points_np=self.mesh_grid_xy,
query_points_np=self.data_grid.over_sampled.array,
)
Expand All @@ -746,6 +816,7 @@ def delaunay(self) -> "scipy.spatial.Delaunay":
split_points=split_points,
splitted_mappings=splitted_mappings,
xp=self._xp,
dual_areas=areas,
)

@property
Expand Down
11 changes: 11 additions & 0 deletions autoarray/inversion/mesh/interpolator/knn.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,17 @@ def kernel_interpolate_points(points, query_chunk, values, k, radius_scale):

class InterpolatorKNearestNeighbor(InterpolatorDelaunay):

@property
def dual_areas(self):
"""
The kNN interpolators build no Delaunay tables, so there are no
in-graph dual areas to hand the mesh geometry. Returning ``None``
keeps ``MeshGeometryDelaunay`` on its host-side fallback (which
triangulates only if the areas are actually asked for) rather than
forcing a qhull call here.
"""
return None

@cached_property
def _mappings_sizes_weights(self):

Expand Down
91 changes: 74 additions & 17 deletions autoarray/inversion/mesh/mesh_geometry/delaunay.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,39 @@ def voronoi_areas_numpy(points, qhull_options="Qbb Qc Qx Qm Q12 Pp"):

class MeshGeometryDelaunay(AbstractMeshGeometry):

def __init__(
self,
mesh,
mesh_grid,
data_grid,
dual_areas=None,
xp=np,
**kwargs,
):
"""
The geometry of a Delaunay triangulation / Voronoi mesh.

Parameters
----------
dual_areas
The barycentric dual area of every mesh vertex, as computed by the
interpolator that built this geometry (see
`autoarray.inversion.mesh.interpolator.delaunay`). These are the
quadrature weights `areas_for_magnification` returns. When the
geometry is built standalone (e.g. in the tests) this is `None` and
the dual areas are computed host-side on demand instead.
"""

super().__init__(
mesh=mesh,
mesh_grid=mesh_grid,
data_grid=data_grid,
xp=xp,
**kwargs,
)

self.dual_areas = dual_areas

@cached_property
def mesh_grid_xy(self):
"""
Expand Down Expand Up @@ -194,23 +227,47 @@ def voronoi_areas(self):
@property
def areas_for_magnification(self) -> np.ndarray:
"""
Returns the Voronoi cell area of every pixel in the mesh, as computed by `voronoi_areas_numpy` (a shoelace
sum over the `scipy.spatial.Voronoi` cell of each mesh point).

Only cells that are **unbounded** in the Voronoi diagram (those `voronoi_areas_numpy` flags with the `-1`
sentinel, because their region runs to infinity and has no finite area) are set to zero. Cells that are
bounded but sit at the edge of the mesh are **kept at full size**, even though they can be far larger than
the interior cells -- an order of magnitude is routine, because a boundary cell extends out to the
circumcentres of the outermost triangles rather than being clipped to the mesh.

These Voronoi areas are **not** the barycentric dual areas used by the Delaunay interpolator
(`barycentric_dual_area_from` in `autoarray.inversion.mesh.interpolator.delaunay`), which assign each vertex
the sum of `triangle_area / 3` over the triangles touching it. The dual areas tile the convex hull of the
mesh exactly and integrate the piecewise-linear reconstruction exactly; these Voronoi areas do neither, so
`sum(reconstruction * areas_for_magnification)` is not the integral of the reconstructed source.
Returns the **barycentric dual area** of every pixel in the mesh: the sum of `triangle_area / 3` over the
Delaunay triangles that touch that vertex.

These are the exact quadrature weights of the mesh's piecewise-linear (barycentric) interpolant. The
Delaunay mapper reconstructs the source as the linear interpolant through the vertex values `s_i`, and for
that interpolant

integral over the hull of f = sum_i s_i * dual_i

holds exactly, not approximately -- the dual areas tile the convex hull of the mesh exactly (they sum to the
hull area), so `sum(reconstruction * areas_for_magnification)` is the integral of the reconstructed source
over the mesh.

The value is taken from the interpolator, which computes it **in-graph** as part of the same Delaunay
construction that builds the mappings (`scipy_delaunay` / `jax_delaunay` and their Matérn variants). It is
therefore free here, and on the JAX path it is a traced value -- so this property is safe to evaluate inside
a `jax.jit` (e.g. PyAutoLens' per-sample latent evaluation). When the geometry is constructed standalone,
without an interpolator (`dual_areas is None`), the dual areas are computed host-side from a
`scipy.spatial.Delaunay` of the mesh grid, which gives the identical answer.

This is **not** `voronoi_areas`, which sums the shoelace area of each point's `scipy.spatial.Voronoi` cell.
Those are correct arithmetic but the wrong quantity for this use: a Voronoi cell that is bounded but sits at
the edge of the mesh extends out to the circumcentres of the outermost triangles instead of being clipped to
the hull, so it can be orders of magnitude larger than that vertex's dual area. Weighting the reconstruction
by them biased magnification by -13% to -99% across the audited configurations (PyAutoArray#522); the fix is
PyAutoArray#524. `voronoi_areas` itself is unchanged and remains available for the geometric uses that
genuinely want a Voronoi cell.
"""
areas = self.voronoi_areas
if self.dual_areas is not None:
return self.dual_areas

areas[areas == -1] = 0.0
import scipy.spatial

from autoarray.inversion.mesh.interpolator.delaunay import (
barycentric_dual_area_from,
)

return areas
mesh_grid_xy = np.asarray(self.mesh_grid_xy)

return barycentric_dual_area_from(
mesh_grid_xy,
scipy.spatial.Delaunay(mesh_grid_xy).simplices,
xp=np,
)
Loading
Loading