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
172 changes: 158 additions & 14 deletions autogalaxy/operate/lens_calc.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
from functools import wraps
import importlib
import logging
import warnings
import numpy as np
from typing import List, Tuple, Union

Expand Down Expand Up @@ -394,8 +395,32 @@ def hessian_from(self, grid, xp=np) -> Tuple:
- **NumPy** (``xp=np``, default): 2-point central finite-difference approximation,
Richardson-extrapolated at step sizes ``h`` and ``h/2`` and combined as
``(4 * H(h/2) - H(h)) / 3``. This cancels the leading ``O(h^2)`` truncation term,
giving ``O(h^4)`` accuracy and matching the JAX path to float64 precision. JAX is
not imported.
giving ``O(h^4)`` accuracy. JAX is not imported.

The step is **adaptive per grid point**, and convergence is judged on the change
between **successive** extrapolants: ``R_k`` from the pair ``(h_k, h_k/2)`` and
``R_{k+1}`` from ``(h_k/2, h_k/4)``, which costs one new finite-difference evaluation
because the half-step evaluation is reused. A point stops refining on either of two
rules:

1. **Tolerance** -- all four Hessian components satisfy
``|R_{k+1} - R_k| <= atol + rtol * |R_{k+1}|`` (defaults ``rtol=1e-7``,
``atol=1e-8``). ``R_{k+1}`` is returned.
2. **Roundoff floor** -- the change has stopped shrinking (``|R_{k+1} - R_k|`` is
larger, relatively, than ``|R_k - R_{k-1}|``) while already below
``roundoff_guard=1e-4`` relative. Halving further only feeds cancellation error, so
``R_k``, the estimate from before the growth, is returned. This is what keeps a
configuration whose deflections are themselves evaluated numerically (a multi-plane
trace, say, which can floor at ~3e-8 relative) from warning on every call.

Neither rule warns. Points that reach ``max_halvings=20`` halvings while their change
is still shrinking -- a genuine singularity, e.g. a grid point exactly on an isothermal
centre, whose relative change never falls below the guard -- keep their final
extrapolant and trigger a single ``UserWarning`` reporting how many points did not
converge. Nothing is ever silently returned as converged, and no exception is raised.

A smooth field settles after three finite-difference evaluations; only points near
compact structure, where a fixed 0.01" step is far too coarse, iterate further.

- **JAX** (``xp=jnp``): exact derivatives via ``jax.jacfwd`` applied to
``deflections_yx_scalar``, vectorised over the grid with ``jnp.vectorize``.
Expand All @@ -414,18 +439,137 @@ def hessian_from(self, grid, xp=np) -> Tuple:
return self._hessian_via_richardson(grid=grid)
return self._hessian_via_jax(grid=grid, xp=xp)

def _hessian_via_richardson(self, grid, buffer: float = 0.01) -> Tuple:
yy_h, xy_h, yx_h, xx_h = self._hessian_via_finite_difference(
grid=grid, buffer=buffer
)
yy_h2, xy_h2, yx_h2, xx_h2 = self._hessian_via_finite_difference(
grid=grid, buffer=buffer / 2.0
)
hessian_yy = (4.0 * yy_h2 - yy_h) / 3.0
hessian_xy = (4.0 * xy_h2 - xy_h) / 3.0
hessian_yx = (4.0 * yx_h2 - yx_h) / 3.0
hessian_xx = (4.0 * xx_h2 - xx_h) / 3.0
return hessian_yy, hessian_xy, hessian_yx, hessian_xx
def _hessian_via_richardson(
self,
grid,
buffer: float = 0.01,
rtol: float = 1.0e-7,
atol: float = 1.0e-8,
max_halvings: int = 20,
roundoff_guard: float = 1.0e-4,
) -> Tuple:
"""
Returns the Hessian via Richardson-extrapolated central finite differences with a step
size that adapts, per grid point, to the scale the deflection field actually varies on.

A finite-difference pair ``H(h)``, ``H(h/2)`` gives the extrapolant
``R = (4 H(h/2) - H(h)) / 3``, which cancels the leading ``O(h^2)`` truncation term. The
fixed-step implementation returned the first such ``R`` whatever the field looked like.
Here the step keeps halving and each point stops on one of two rules, neither of which
warns:

1. **Tolerance.** ``R_k`` from ``(h_k, h_k/2)`` and ``R_{k+1}`` from ``(h_k/2, h_k/4)``
agree on all four components: ``|R_{k+1} - R_k| <= atol + rtol * |R_{k+1}|``.
``R_{k+1}`` is kept.
2. **Roundoff floor.** The relative change grew instead of shrinking, while already
below ``roundoff_guard``. Finite differences fall as ``O(h^2)`` only until
cancellation in ``(f(x+h) - f(x-h))`` takes over, after which halving makes the answer
*worse*; the turning point is the best the arithmetic can do. ``R_k``, the estimate
from before the growth, is kept.

Because the half-step evaluation is reused as the next full-step one, each halving costs a
single finite-difference evaluation, and only on the shrinking active subset.

Judging on the extrapolants rather than on a pair's own error estimate
``|H(h/2) - H(h)| / 3`` matters: that estimate bounds the error of ``H(h/2)``, which is
``O(h^2)``, not of ``R``, which is ``O(h^4)``, so it declares non-convergence orders of
magnitude past the point where the returned value has stopped moving.

Rule 2 is guarded by ``roundoff_guard`` so that it cannot silence a genuine singularity:
a grid point sitting exactly on an isothermal centre has a relative change that stays at
~0.5 forever, never entering the guard band, so it runs to ``max_halvings``, keeps its
last extrapolant and raises a single ``UserWarning``. It is never silently accepted, and
no exception is raised (a raise here would kill an otherwise-converged model fit).

Parameters
----------
grid
The 2D grid of (y,x) arc-second coordinates the Hessian is computed on.
buffer
The initial finite-difference step size in arc-seconds.
rtol
The relative tolerance the change between successive extrapolants must meet.
atol
The absolute tolerance the change between successive extrapolants must meet.
max_halvings
The maximum number of times the step is halved before points that are still refining
are warned about and their last value kept.
roundoff_guard
The relative change below which a growing change is read as the roundoff floor rather
than as a field the step has not resolved yet.
"""
grid_values = grid.array if hasattr(grid, "array") else grid
grid_values = np.array(grid_values, dtype=np.float64, copy=True)

hessian_full = np.stack(
self._hessian_via_finite_difference(grid=grid, buffer=buffer)
).astype(np.float64)
hessian_half = np.stack(
self._hessian_via_finite_difference(grid=grid, buffer=buffer / 2.0)
).astype(np.float64)

richardson = (4.0 * hessian_half - hessian_full) / 3.0

total = richardson.shape[1]
relative_change = np.full(total, np.inf)
refining = np.ones(total, dtype=bool)

step = buffer
halvings = 0

while np.any(refining) and halvings < max_halvings:
halvings += 1
step /= 2.0

index = np.flatnonzero(refining)

grid_subset = aa.Grid2DIrregular(values=grid_values[index])

full_subset = hessian_half[:, index]
half_subset = np.stack(
self._hessian_via_finite_difference(grid=grid_subset, buffer=step / 2.0)
).astype(np.float64)

richardson_subset = (4.0 * half_subset - full_subset) / 3.0

change = np.abs(richardson_subset - richardson[:, index])

within_tolerance = np.all(
change <= atol + rtol * np.abs(richardson_subset), axis=0
)

denominator = np.where(
np.abs(richardson_subset) > 0.0, np.abs(richardson_subset), 1.0
)
change_relative = np.max(change / denominator, axis=0)

at_roundoff_floor = (change_relative > relative_change[index]) & (
relative_change[index] <= roundoff_guard
)

# Points at the roundoff floor keep R_k, the extrapolant from before the change
# started growing, so their columns are left untouched; every other point adopts
# the new extrapolant and the half-step evaluation that will seed its next pair.
adopted = index[~at_roundoff_floor]

richardson[:, adopted] = richardson_subset[:, ~at_roundoff_floor]
hessian_half[:, adopted] = half_subset[:, ~at_roundoff_floor]
relative_change[adopted] = change_relative[~at_roundoff_floor]

refining = np.zeros(total, dtype=bool)
refining[index] = ~(within_tolerance | at_roundoff_floor)

if np.any(refining):
number = int(np.count_nonzero(refining))
largest = float(np.max(relative_change[refining]))
warnings.warn(
f"LensCalc Hessian: {number} of {total} points did not converge after "
f"{max_halvings} halvings (largest relative error estimate "
f"{largest:.2e}); values kept.",
UserWarning,
)

return richardson[0], richardson[1], richardson[2], richardson[3]

def _hessian_via_jax(self, grid, xp) -> Tuple:
import jax
Expand Down
129 changes: 129 additions & 0 deletions test_autogalaxy/operate/test_deflections.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,6 +129,135 @@ def test__hessian_from__axis_aligned_grid__correct_values():
assert hessian_xx == pytest.approx(np.array([2.22209, 0.0]), 1.0e-4)


def test__hessian_from__adaptive_step__compact_sis_near_centre():
"""
Regression test for the adaptive Richardson step (issue #591).

Close to the centre of a compact deflector the deflection field varies on a scale far smaller
than the 0.01" step the NumPy Hessian used to be hardcoded to, so the finite differences
straddled the whole deflector and returned values that were ~100% wrong (and, for the
magnification, sign-flipped). With the step adapted per point the Hessian must reproduce the
profile's own analytic shear and convergence.

``IsothermalSph`` has closed-form shear and convergence, so the analytic values are an
independent oracle: at radius ``r`` from the centre both have magnitude
``einstein_radius / (2 r)``, i.e. of order 100-330 for the radii used here.
"""
mp = ag.mp.IsothermalSph(centre=(0.0, 0.0), einstein_radius=0.2)

radii = [3.0e-4, 4.0e-4, 5.5e-4, 7.0e-4, 8.5e-4, 1.0e-3]
angles = [10.0, 55.0, 100.0, 170.0, 230.0, 310.0]

grid = ag.Grid2DIrregular(
values=[
(
radius * math.sin(math.radians(angle)),
radius * math.cos(math.radians(angle)),
)
for radius, angle in zip(radii, angles)
]
)

od = LensCalc.from_mass_obj(mp)

shear_analytic = np.asarray(mp.shear_yx_2d_from(grid=grid))
shear_via_hessian = np.asarray(od.shear_yx_2d_via_hessian_from(grid=grid))

convergence_analytic = np.asarray(mp.convergence_2d_from(grid=grid))
convergence_via_hessian = np.asarray(od.convergence_2d_via_hessian_from(grid=grid))

np.testing.assert_allclose(shear_via_hessian, shear_analytic, rtol=1.0e-4)
np.testing.assert_allclose(
convergence_via_hessian, convergence_analytic, rtol=1.0e-4
)


def test__hessian_from__adaptive_step__smooth_field_unchanged():
"""
The adaptive step of issue #591 must not move the answer on the smooth fields the fixed-step
implementation already handled well: the values pinned by
``test__hessian_from__diagonal_grid__correct_values`` and
``test__hessian_from__axis_aligned_grid__correct_values`` must still hold, and the adaptive
result must agree with the old fixed-step Richardson extrapolation (one pair at h=0.01 and
h=0.005) to far inside those tests' tolerances.

Where the two differ at all it is the fixed step's own residual truncation error: the adaptive
result is the more accurate of the two (checked here against the profile's analytic
convergence, which it reproduces to ~1e-12 against the fixed step's ~1e-9).
"""
mp = ag.mp.Isothermal(
centre=(0.0, 0.0), ell_comps=(0.0, -0.111111), einstein_radius=2.0
)

od = LensCalc.from_mass_obj(mp)

grid_diagonal = ag.Grid2DIrregular(values=[(0.5, 0.5), (1.0, 1.0)])
grid_axis_aligned = ag.Grid2DIrregular(values=[(1.0, 0.0), (0.0, 1.0)])

for grid in [grid_diagonal, grid_axis_aligned]:
hessian_h = np.stack(od._hessian_via_finite_difference(grid=grid, buffer=0.01))
hessian_h2 = np.stack(
od._hessian_via_finite_difference(grid=grid, buffer=0.005)
)
hessian_fixed_step = (4.0 * hessian_h2 - hessian_h) / 3.0

hessian_adaptive = np.stack(od.hessian_from(grid=grid))

np.testing.assert_allclose(
hessian_adaptive, hessian_fixed_step, rtol=1.0e-7, atol=1.0e-10
)

convergence_adaptive = 0.5 * (hessian_adaptive[0] + hessian_adaptive[3])
convergence_fixed_step = 0.5 * (hessian_fixed_step[0] + hessian_fixed_step[3])
convergence_analytic = np.asarray(mp.convergence_2d_from(grid=grid))

assert np.max(np.abs(convergence_adaptive - convergence_analytic)) <= np.max(
np.abs(convergence_fixed_step - convergence_analytic)
)

hessian_yy, hessian_xy, hessian_yx, hessian_xx = od.hessian_from(grid=grid_diagonal)

assert hessian_yy == pytest.approx(np.array([1.3882113, 0.6941056]), 1.0e-4)
assert hessian_xy == pytest.approx(np.array([-1.3882113, -0.6941056]), 1.0e-4)
assert hessian_yx == pytest.approx(np.array([-1.3882113, -0.6941056]), 1.0e-4)
assert hessian_xx == pytest.approx(np.array([1.3882113, 0.6941056]), 1.0e-4)

hessian_yy, hessian_xy, hessian_yx, hessian_xx = od.hessian_from(
grid=grid_axis_aligned
)

assert hessian_yy == pytest.approx(np.array([0.0, 1.777699]), 1.0e-4)
assert hessian_xy == pytest.approx(np.array([0.0, 0.0]), 1.0e-4)
assert hessian_yx == pytest.approx(np.array([0.0, 0.0]), 1.0e-4)
assert hessian_xx == pytest.approx(np.array([2.22209, 0.0]), 1.0e-4)


def test__hessian_from__unconverged_points_warn():
"""
The deflection angles of an isothermal sphere are discontinuous at its centre, so no step size
resolves the Hessian there and the adaptive refinement must give up. It must do so loudly (a
single ``UserWarning``), never silently and never by raising -- a raise here would kill an
otherwise-converged model fit. The values it keeps must stay finite, and the smooth point on
the same grid, which converged on the first pair, must be unaffected.
"""
mp = ag.mp.IsothermalSph(centre=(0.0, 0.0), einstein_radius=1.0)

grid = ag.Grid2DIrregular(values=[(0.0, 0.0), (1.0, 1.0)])

od = LensCalc.from_mass_obj(mp)

with pytest.warns(UserWarning, match="did not converge"):
hessian = od.hessian_from(grid=grid)

assert np.all(np.isfinite(np.stack(hessian)))

convergence_smooth = 0.5 * (hessian[0][1] + hessian[3][1])

assert convergence_smooth == pytest.approx(
float(np.asarray(mp.convergence_2d_from(grid=grid))[1]), 1.0e-4
)


def test__convergence_2d_via_hessian_from():
grid = ag.Grid2DIrregular(
values=[(1.075, -0.125), (-0.875, -0.075), (-0.925, -0.075), (0.075, 0.925)]
Expand Down
Loading