From b91b332f7e30597051ee6bad8217ca08bc6f0052 Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sat, 29 Aug 2026 16:08:26 -0400 Subject: [PATCH] =?UTF-8?q?feat:=20AdaptPower=20regularization=20siblings?= =?UTF-8?q?=20=E2=80=94=20power=20input=20+=20single-scatter=20fix?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds AdaptPower, AdaptSplitPower, AdaptSplitZerothPower and MaternAdaptPowerKernel as corrected siblings of the Adapt family. The legacy classes and their util functions are untouched, so stored af.Model identifiers, output directories and aggregator reloads stay valid. Two things separate the new classes from the legacy ones: 1. The coefficient enters the regularization matrix at 2 * power (default power=1.0, the Constant lambda^2 convention) rather than always at lambda^4. adapt_regularization_weights_from gains a power keyword defaulting to 2.0, so the legacy classes keep their exact numerics. 2. weighted_regularization_matrix_single_scatter_from scatters each mesh edge once, so AdaptPower(inner=outer=c) equals Constant(c) exactly instead of being 2x it. The split family already shared its builder with ConstantSplit and needed no change. Motivation: under the shared LogUniform(1e-6, 1e6) prior, the lambda^4 scale drives the regularization matrix non-positive-definite from c ~ 1e4 rather than c ~ 1e6, which is the mechanism behind the likelihood-overflow flood seen on RAL pilot 341908_5. Migration: c_new = c_old ** 2; power=2.0 restores the legacy coefficient scaling. Refs #511 --- .../inversion/regularization/__init__.py | 4 + autoarray/inversion/regularization/adapt.py | 123 +++++++++- .../inversion/regularization/adapt_power.py | 137 ++++++++++++ .../inversion/regularization/adapt_split.py | 21 +- .../regularization/adapt_split_power.py | 107 +++++++++ .../regularization/adapt_split_zeroth.py | 18 ++ .../adapt_split_zeroth_power.py | 109 +++++++++ .../regularization/matern_adapt_kernel.py | 21 +- .../matern_adapt_power_kernel.py | 115 ++++++++++ .../regularization/regularization_util.py | 3 + .../regularizations/test_adapt_power.py | 211 ++++++++++++++++++ .../regularizations/test_adapt_power_jax.py | 126 +++++++++++ .../regularizations/test_adapt_split_power.py | 138 ++++++++++++ 13 files changed, 1128 insertions(+), 5 deletions(-) create mode 100644 autoarray/inversion/regularization/adapt_power.py create mode 100644 autoarray/inversion/regularization/adapt_split_power.py create mode 100644 autoarray/inversion/regularization/adapt_split_zeroth_power.py create mode 100644 autoarray/inversion/regularization/matern_adapt_power_kernel.py create mode 100644 test_autoarray/inversion/regularizations/test_adapt_power.py create mode 100644 test_autoarray/inversion/regularizations/test_adapt_power_jax.py create mode 100644 test_autoarray/inversion/regularizations/test_adapt_split_power.py diff --git a/autoarray/inversion/regularization/__init__.py b/autoarray/inversion/regularization/__init__.py index f8a3160cf..5bcc2d5b0 100644 --- a/autoarray/inversion/regularization/__init__.py +++ b/autoarray/inversion/regularization/__init__.py @@ -4,12 +4,16 @@ from .constant_zeroth import ConstantZeroth from .constant_split import ConstantSplit from .adapt import Adapt +from .adapt_power import AdaptPower from .adapt_split import AdaptSplit +from .adapt_split_power import AdaptSplitPower from .brightness_zeroth import BrightnessZeroth from .adapt_split_zeroth import AdaptSplitZeroth +from .adapt_split_zeroth_power import AdaptSplitZerothPower from .curvature_mask import CurvatureMask from .fourth_order_mask import FourthOrderMask from .gaussian_kernel import GaussianKernel from .exponential_kernel import ExponentialKernel from .matern_kernel import MaternKernel from .matern_adapt_kernel import MaternAdaptKernel +from .matern_adapt_power_kernel import MaternAdaptPowerKernel diff --git a/autoarray/inversion/regularization/adapt.py b/autoarray/inversion/regularization/adapt.py index 4c8b84a0f..ec3d36dc9 100644 --- a/autoarray/inversion/regularization/adapt.py +++ b/autoarray/inversion/regularization/adapt.py @@ -10,7 +10,10 @@ def adapt_regularization_weights_from( - inner_coefficient: float, outer_coefficient: float, pixel_signals: np.ndarray + inner_coefficient: float, + outer_coefficient: float, + pixel_signals: np.ndarray, + power: float = 2.0, ) -> np.ndarray: """ Returns the regularization weights for the adaptive regularization scheme (e.g. ``Adapt``). @@ -36,6 +39,14 @@ def adapt_regularization_weights_from( pixel_signals The estimated signal in every pixelization pixel, used to change the regularization weighting of high signal and low signal pixelizations. + power + The exponent the interpolated coefficient is raised to. The matrix builders square the returned weights + again, so the coefficient enters the regularization matrix at the power ``2 * power``. + + The default ``2.0`` is the historical ``Adapt`` convention (a fourth-power coefficient dependence) and is + what the legacy ``Adapt``, ``AdaptSplit``, ``AdaptSplitZeroth`` and ``MaternAdaptKernel`` classes pass. + ``power=1.0`` gives the squared-once convention shared with ``Constant`` and is what the ``*Power`` + classes (e.g. ``AdaptPower``) pass by default. Returns ------- @@ -45,7 +56,7 @@ def adapt_regularization_weights_from( """ return ( inner_coefficient * pixel_signals + outer_coefficient * (1.0 - pixel_signals) - ) ** 2.0 + ) ** power def weighted_regularization_matrix_from( @@ -129,6 +140,91 @@ def weighted_regularization_matrix_from( return mat[:S, :S] +def weighted_regularization_matrix_single_scatter_from( + regularization_weights: np.ndarray, + neighbors: np.ndarray, + xp=np, +) -> np.ndarray: + """ + Returns the regularization matrix of the adaptive regularization scheme, scattering every mesh edge + **once** so that uniform weights reproduce ``Constant`` regularization exactly. + + This is the corrected sibling of ``weighted_regularization_matrix_from``, used by the ``*Power`` + classes (e.g. ``AdaptPower``). The two differ only in how often each mesh edge is scattered: + + - ``weighted_regularization_matrix_from`` adds every ordered neighbor pair to **both** the + ``(i, j)`` and ``(j, i)`` entries. Because the neighbor list already holds each unordered edge + twice (once in row ``i``, once in row ``j``), every edge lands four times, which is exactly twice + what ``constant_regularization_matrix_from`` does. ``Adapt(inner=outer=c)`` is therefore + ``2 x`` ``Constant(c)``, not equal to it. + - This function adds every ordered neighbor pair once, to ``(i, i)`` and ``(i, j)`` only -- the same + bookkeeping ``Constant`` uses -- so ``AdaptPower(inner=outer=c)`` equals ``Constant(c)`` exactly. + + The edge weight is the mean of the two endpoints' squared regularization weights, + ``0.5 * (w_i ** 2 + w_j ** 2)``. This is symmetric in ``(i, j)`` (so the matrix is symmetric even + for wildly varying adaptive weights), reduces to ``w ** 2`` when the weights are uniform, and makes + the result a weighted graph Laplacian plus a ``1e-8`` diagonal floor -- so its rows sum to the floor + and it is positive semi-definite by construction. + + The legacy builder is left untouched: halving it in place would silently change the effective + regularization of every ``Adapt`` fit ever run. + + Parameters + ---------- + regularization_weights + The regularization weight of each pixel, adaptively governing the degree of gradient regularization + applied to each inversion parameter (e.g. mesh pixels of a ``Mapper``). + neighbors + An array of length (total_pixels) which provides the index of all neighbors of every pixel in + the mesh grid (entries of -1 correspond to no neighbor). + + Returns + ------- + np.ndarray + The regularization matrix computed using an adaptive regularization scheme where the effective + regularization coefficient of every source pixel is different. + """ + S, P = neighbors.shape + + reg_w = regularization_weights**2 + + # 1) Flatten the (i->j) neighbor pairs + I = xp.repeat(xp.arange(S), P) # (S*P,) + J_raw = neighbors.reshape(-1) # (S*P,) + + # 2) Remap "no neighbor" entries to an extra slot S, whose weight = 0 + OUT = S + valid = J_raw >= 0 + J = xp.where(valid, J_raw, OUT) + + # 3) Build an extended weight vector with a zero at index S + reg_w_ext = xp.concatenate([reg_w, xp.zeros((1,))], axis=0) + + # 4) Symmetric per-edge weight: the mean of the two endpoints' squared weights. Padded entries are + # masked to zero (their I endpoint is a real pixel, so the mean alone would not vanish). + w_ij = 0.5 * (reg_w_ext[I] + reg_w_ext[J]) * valid + + # 5) Start with zeros on an (S+1)x(S+1) canvas so we can scatter into row S safely + mat = xp.zeros((S + 1, S + 1), dtype=regularization_weights.dtype) + + diag_updates_i = xp.concatenate( + [xp.full((S,), 1e-8), xp.zeros((1,))], axis=0 # out-of-bounds slot stays zero + ) + + # 6) Scatter each ordered pair exactly once: onto the diagonal of i and the (i, j) off-diagonal + if xp.__name__.startswith("jax"): + mat = mat.at[xp.diag_indices(S + 1)].add(diag_updates_i) + mat = mat.at[I, I].add(w_ij) + mat = mat.at[I, J].add(-w_ij) + else: + np.add.at(mat, np.diag_indices(S + 1), diag_updates_i) + np.add.at(mat, (I, I), w_ij) + np.add.at(mat, (I, J), -w_ij) + + # 7) Drop the extra row/column S and return the SxS result + return mat[:S, :S] + + class Adapt(AbstractRegularization): def __init__( self, @@ -180,7 +276,28 @@ def __init__( neighbors come from a direct scipy call on the traced mesh grid (use ``AdaptSplit`` there). Note the defaults ``inner_coefficient == outer_coefficient == 1.0`` make the weighting - uniform — numerically identical to ``Constant(coefficient=1.0)``. + uniform — but *not* numerically identical to ``Constant(coefficient=1.0)``; + see the coefficient-convention note below. + + **Coefficient convention (legacy, ``lambda^4``).** The coefficients are squared twice before they + reach the regularization matrix -- once by ``adapt_regularization_weights_from`` and once by the + matrix builder -- so the matrix scales as the *fourth* power of the coefficient, while + ``Constant`` scales as the second. Both carry the same ``LogUniform(1e-6, 1e6)`` prior, so this + scheme explores a far wider effective smoothing range and reaches a numerically non + positive-definite matrix from ``c ~ 1e4`` where ``Constant`` survives to ``c ~ 1e6``. + + **It is also 2x ``Constant``, not equal to it.** The matrix builder scatters every mesh edge in + both directions, and the neighbor list already holds each unordered edge twice, so each edge + lands four times where ``Constant`` lands it twice. + ``Adapt(inner_coefficient=1.0, outer_coefficient=1.0)`` is therefore exactly ``2 x`` + ``Constant(coefficient=1.0)``. + + This behaviour is preserved deliberately: changing it would alter the coefficient scale of every + adaptive fit ever run. **New work should prefer ``AdaptPower``**, which takes a ``power`` argument + (default ``1.0``, giving the ``Constant``-matching ``lambda^2`` convention) and scatters each edge + once, so ``AdaptPower(inner=outer=c)`` equals ``Constant(c)`` exactly and is more robust to + gradient / NaN pathologies. The migration is ``c_new = c_old ** 2``, and + ``AdaptPower(power=2.0)`` reproduces this class's coefficient scaling exactly. Parameters ---------- diff --git a/autoarray/inversion/regularization/adapt_power.py b/autoarray/inversion/regularization/adapt_power.py new file mode 100644 index 000000000..6d4a84417 --- /dev/null +++ b/autoarray/inversion/regularization/adapt_power.py @@ -0,0 +1,137 @@ +from __future__ import annotations +import numpy as np +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from autoarray.inversion.linear_obj.linear_obj import LinearObj + +from autoarray.inversion.regularization.adapt import Adapt +from autoarray.inversion.regularization.adapt import adapt_regularization_weights_from +from autoarray.inversion.regularization.adapt import ( + weighted_regularization_matrix_single_scatter_from, +) + + +class AdaptPower(Adapt): + def __init__( + self, + inner_coefficient: float = 1.0, + outer_coefficient: float = 1.0, + signal_scale: float = 1.0, + power: float = 1.0, + ): + """ + Regularization which uses the neighbors of the mesh (e.g. shared Delaunay vertexes) and values adapted to + the data being fitted to smooth an inversion's solution, with the coefficient convention of ``Constant``. + + This is the corrected sibling of ``Adapt``. It reconstructs the same adaptive weighting -- high smoothing + where there is no signal, less where there is (Nightingale, Dye and Massey 2018) -- but fixes the two ways + in which ``Adapt`` diverges from ``Constant``: + + 1. **The coefficient enters at ``lambda^2``, not ``lambda^4``.** ``Adapt`` squares its coefficients twice + (once when interpolating them into per-pixel weights, once in the matrix builder), so its matrix scales + as the fourth power of the coefficient. This class raises the interpolated coefficient to ``power`` + before the builder squares it, so the effective exponent is ``2 * power`` and the default ``power=1.0`` + matches ``Constant``. Under the shared ``LogUniform(1e-6, 1e6)`` prior that means the prior now spans + ``lambda^2``, and the regularization matrix stays positive-definite to ``c ~ 1e6`` rather than + collapsing from ``c ~ 1e4`` -- the fragility that produced the likelihood-overflow floods seen in + free-coefficient adaptive fits. + 2. **Every mesh edge is scattered once.** ``Adapt`` adds each ordered neighbor pair to both the ``(i, j)`` + and ``(j, i)`` entries, and the neighbor list already holds each unordered edge twice, so its matrix is + exactly ``2 x`` ``Constant``'s. This class uses + ``weighted_regularization_matrix_single_scatter_from``, which scatters each ordered pair once with the + symmetric edge weight ``0.5 * (w_i ** 2 + w_j ** 2)`` -- a weighted graph Laplacian, so still symmetric + and positive semi-definite. + + Together these make ``AdaptPower(inner_coefficient=c, outer_coefficient=c)`` **exactly equal** to + ``Constant(coefficient=c)`` for any ``c``, which is what ``Adapt``'s docstring always claimed and never + delivered. + + A full description of regularization and this matrix can be found in the parent ``AbstractRegularization`` + class; the ``B`` matrix construction is described on ``Adapt``. + + **Migration from ``Adapt``.** The coefficient scale is squared: ``c_new = c_old ** 2``. To reproduce the + legacy class exactly, pass ``power=2.0`` -- but note the factor-2 scatter is fixed regardless, so + ``AdaptPower(power=2.0)`` is ``0.5 x`` ``Adapt`` with the same coefficients. + + **JAX & gradient support**: as for ``Adapt`` -- JAX-differentiable and FD-certified on the rectangular + mesh family, but raising ``TracerArrayConversionError`` on the Delaunay mesh family, whose neighbors come + from a direct scipy call on the traced mesh grid (use ``AdaptSplitPower`` there). + + Parameters + ---------- + inner_coefficient + The inner regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the inner (high signal) regions of a mesh's reconstruction. + outer_coefficient + The outer regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the outer (low signal) regions of a mesh's reconstruction. + signal_scale + A factor which controls how rapidly the smoothness of regularization varies from high signal regions + to low signal regions. + power + The exponent the interpolated coefficient is raised to before the matrix builder squares it, so the + coefficient enters the regularization matrix at the power ``2 * power``. The default ``1.0`` is the + ``Constant`` convention; ``2.0`` is the legacy ``Adapt`` convention. This is a convention switch, not + a model parameter -- the shipped prior config fixes it as a ``Constant`` prior so a search never + samples it. + """ + super().__init__( + inner_coefficient=inner_coefficient, + outer_coefficient=outer_coefficient, + signal_scale=signal_scale, + ) + + self.power = power + + def regularization_weights_from(self, linear_obj: LinearObj, xp=np) -> np.ndarray: + """ + Returns the regularization weights of this regularization scheme. + + These are the interpolated inner / outer coefficients raised to ``self.power`` (default ``1.0``), as + opposed to ``Adapt``, which squares them. + + Parameters + ---------- + linear_obj + The linear object (e.g. a ``Mapper``) which uses these weights when performing regularization. + + Returns + ------- + The regularization weights. + """ + pixel_signals = linear_obj.pixel_signals_from( + signal_scale=self.signal_scale, xp=xp + ) + + return adapt_regularization_weights_from( + inner_coefficient=self.inner_coefficient, + outer_coefficient=self.outer_coefficient, + pixel_signals=pixel_signals, + power=self.power, + ) + + def regularization_matrix_from(self, linear_obj: LinearObj, xp=np) -> np.ndarray: + """ + Returns the regularization matrix with shape [pixels, pixels]. + + Every mesh edge is scattered once (unlike ``Adapt``), so uniform weights reproduce ``Constant`` exactly. + + Parameters + ---------- + linear_obj + The linear object (e.g. a ``Mapper``) which uses this matrix to perform regularization. + + Returns + ------- + The regularization matrix. + """ + regularization_weights = self.regularization_weights_from( + linear_obj=linear_obj, xp=xp + ) + + return weighted_regularization_matrix_single_scatter_from( + regularization_weights=regularization_weights, + neighbors=linear_obj.mesh_geometry.neighbors, + xp=xp, + ) diff --git a/autoarray/inversion/regularization/adapt_split.py b/autoarray/inversion/regularization/adapt_split.py index dd2bf6b77..8078efc30 100644 --- a/autoarray/inversion/regularization/adapt_split.py +++ b/autoarray/inversion/regularization/adapt_split.py @@ -67,7 +67,26 @@ def __init__( mesh family (e.g. the KNN meshes), structurally incompatible with the rectangular meshes. Note the defaults ``inner_coefficient == outer_coefficient == 1.0`` make the weighting - uniform — numerically identical to ``ConstantSplit(coefficient=1.0)``. + uniform — but *not* numerically identical to ``ConstantSplit(coefficient=1.0)``; + see the coefficient-convention note below. + + **Coefficient convention (legacy, ``lambda^4``).** The coefficients are squared twice before they + reach the regularization matrix -- once by ``adapt_regularization_weights_from`` and once by the + matrix builder -- so the matrix scales as the *fourth* power of the coefficient, while + ``Constant`` scales as the second. Both carry the same ``LogUniform(1e-6, 1e6)`` prior, so this + scheme explores a far wider effective smoothing range and reaches a numerically non + positive-definite matrix from ``c ~ 1e4`` where ``Constant`` survives to ``c ~ 1e6``. + + The split family does **not** carry the factor-2 scatter asymmetry of ``Adapt``: + it shares ``pixel_splitted_regularization_matrix_from`` with ``ConstantSplit``, so the + coefficient exponent is the only difference between the two. + + This behaviour is preserved deliberately: changing it would alter the coefficient scale of every + adaptive fit ever run. **New work should prefer ``AdaptSplitPower``**, which takes a ``power`` argument + (default ``1.0``, giving the ``Constant``-matching ``lambda^2`` convention), so + ``AdaptSplitPower(inner=outer=c)`` equals ``ConstantSplit(c)`` exactly and is more robust to + gradient / NaN pathologies. The migration is ``c_new = c_old ** 2``, and + ``AdaptSplitPower(power=2.0)`` reproduces this class's coefficient scaling exactly. Parameters ---------- diff --git a/autoarray/inversion/regularization/adapt_split_power.py b/autoarray/inversion/regularization/adapt_split_power.py new file mode 100644 index 000000000..bb6c3b468 --- /dev/null +++ b/autoarray/inversion/regularization/adapt_split_power.py @@ -0,0 +1,107 @@ +from __future__ import annotations +import numpy as np +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from autoarray.inversion.linear_obj.linear_obj import LinearObj + +from autoarray.inversion.regularization.adapt import adapt_regularization_weights_from +from autoarray.inversion.regularization.adapt_split import AdaptSplit + + +class AdaptSplitPower(AdaptSplit): + is_split_regularization = True + + def __init__( + self, + inner_coefficient: float = 1.0, + outer_coefficient: float = 1.0, + signal_scale: float = 1.0, + power: float = 1.0, + ): + """ + Regularization which uses the derivatives at a cross of four points around each pixel centre and values + adapted to the data being fitted to smooth an inversion's solution, with the coefficient convention of + ``ConstantSplit``. + + This is the corrected sibling of ``AdaptSplit``. The split geometry, the interpolation to the cross of + four regularization points and the matrix builder are all unchanged -- the only difference is the + coefficient convention: + + ``AdaptSplit`` squares its coefficients twice (once when interpolating them into per-pixel weights, once + in the matrix builder), so its matrix scales as the fourth power of the coefficient while + ``ConstantSplit`` scales as the second. This class raises the interpolated coefficient to ``power`` + before the builder squares it, so the effective exponent is ``2 * power`` and the default ``power=1.0`` + matches ``ConstantSplit``. Under the shared ``LogUniform(1e-6, 1e6)`` prior the prior now spans + ``lambda^2``, and the regularization matrix stays positive-definite to ``c ~ 1e6`` rather than collapsing + from ``c ~ 1e4`` -- the fragility that produced the likelihood-overflow floods seen in free-coefficient + adaptive fits. + + This makes ``AdaptSplitPower(inner_coefficient=c, outer_coefficient=c)`` **exactly equal** to + ``ConstantSplit(coefficient=c)`` for any ``c``. + + The split family never carried ``Adapt``'s factor-2 scatter asymmetry: ``AdaptSplit`` and + ``ConstantSplit`` already share ``pixel_splitted_regularization_matrix_from``, which scatters each + contribution once. + + A visual description of the split scheme is in the appendix of He et al. (2024): + https://arxiv.org/abs/2403.16253 + + **Migration from ``AdaptSplit``.** The coefficient scale is squared: ``c_new = c_old ** 2``. To reproduce + the legacy class exactly, pass ``power=2.0``. + + **JAX & gradient support**: as for ``AdaptSplit`` -- differentiable and FD-certified on the Delaunay mesh + family (e.g. the KNN meshes), structurally incompatible with the rectangular meshes. + + Parameters + ---------- + inner_coefficient + The inner regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the inner (high signal) regions of a mesh's reconstruction. + outer_coefficient + The outer regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the outer (low signal) regions of a mesh's reconstruction. + signal_scale + A factor which controls how rapidly the smoothness of regularization varies from high signal regions + to low signal regions. + power + The exponent the interpolated coefficient is raised to before the matrix builder squares it, so the + coefficient enters the regularization matrix at the power ``2 * power``. The default ``1.0`` is the + ``ConstantSplit`` convention; ``2.0`` is the legacy ``AdaptSplit`` convention. This is a convention + switch, not a model parameter -- the shipped prior config fixes it as a ``Constant`` prior so a + search never samples it. + """ + super().__init__( + inner_coefficient=inner_coefficient, + outer_coefficient=outer_coefficient, + signal_scale=signal_scale, + ) + + self.power = power + + def regularization_weights_from(self, linear_obj: LinearObj, xp=np) -> np.ndarray: + """ + Returns the regularization weights of this regularization scheme. + + These are the interpolated inner / outer coefficients raised to ``self.power`` (default ``1.0``), as + opposed to ``AdaptSplit``, which squares them. + + Parameters + ---------- + linear_obj + The linear object (e.g. a ``Mapper``) which uses these weights when performing regularization. + + Returns + ------- + The regularization weights. + """ + pixel_signals = linear_obj.pixel_signals_from( + signal_scale=self.signal_scale, xp=xp + ) + + return adapt_regularization_weights_from( + inner_coefficient=self.inner_coefficient, + outer_coefficient=self.outer_coefficient, + pixel_signals=pixel_signals, + power=self.power, + ) diff --git a/autoarray/inversion/regularization/adapt_split_zeroth.py b/autoarray/inversion/regularization/adapt_split_zeroth.py index 670392cbb..f6d151183 100644 --- a/autoarray/inversion/regularization/adapt_split_zeroth.py +++ b/autoarray/inversion/regularization/adapt_split_zeroth.py @@ -69,6 +69,24 @@ def __init__( blocker. Still structurally incompatible with the rectangular meshes (the split leg's shape mismatch, as ``ConstantSplit``). + **Coefficient convention (legacy, ``lambda^4``).** The coefficients are squared twice before they + reach the regularization matrix -- once by ``adapt_regularization_weights_from`` and once by the + matrix builder -- so the matrix scales as the *fourth* power of the coefficient, while + ``Constant`` scales as the second. Both carry the same ``LogUniform(1e-6, 1e6)`` prior, so this + scheme explores a far wider effective smoothing range and reaches a numerically non + positive-definite matrix from ``c ~ 1e4`` where ``Constant`` survives to ``c ~ 1e6``. + + The split family does **not** carry the factor-2 scatter asymmetry of ``Adapt``: + it shares ``pixel_splitted_regularization_matrix_from`` with ``ConstantSplit``, so the + coefficient exponent is the only difference between the two. The zeroth leg (``BrightnessZeroth``) is + already squared once and is unaffected. + + This behaviour is preserved deliberately: changing it would alter the coefficient scale of every + adaptive fit ever run. **New work should prefer ``AdaptSplitZerothPower``**, which takes a ``power`` argument + (default ``1.0``, giving the ``Constant``-matching ``lambda^2`` convention) and is more robust to + gradient / NaN pathologies. The migration is ``c_new = c_old ** 2``, and + ``AdaptSplitZerothPower(power=2.0)`` reproduces this class's coefficient scaling exactly. + Parameters ---------- coefficients diff --git a/autoarray/inversion/regularization/adapt_split_zeroth_power.py b/autoarray/inversion/regularization/adapt_split_zeroth_power.py new file mode 100644 index 000000000..b7b2abadb --- /dev/null +++ b/autoarray/inversion/regularization/adapt_split_zeroth_power.py @@ -0,0 +1,109 @@ +from __future__ import annotations +import numpy as np +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from autoarray.inversion.linear_obj.linear_obj import LinearObj + +from autoarray.inversion.regularization.adapt import adapt_regularization_weights_from +from autoarray.inversion.regularization.adapt_split_zeroth import AdaptSplitZeroth + + +class AdaptSplitZerothPower(AdaptSplitZeroth): + is_split_regularization = True + + def __init__( + self, + zeroth_coefficient: float = 1.0, + zeroth_signal_scale: float = 1.0, + inner_coefficient: float = 1.0, + outer_coefficient: float = 1.0, + signal_scale: float = 1.0, + power: float = 1.0, + ): + """ + Regularization which uses zeroth order regularization, the derivatives at a cross of four points around + each pixel centre and values adapted to the data being fitted to smooth an inversion's solution, with the + coefficient convention of ``ConstantSplit``. + + This is the corrected sibling of ``AdaptSplitZeroth``. The split geometry, the zeroth-order leg and the + matrix builders are all unchanged -- the only difference is the convention of the split leg's + ``inner_coefficient`` / ``outer_coefficient``: + + ``AdaptSplitZeroth`` squares them twice (once when interpolating them into per-pixel weights, once in the + matrix builder), so the split leg scales as the fourth power of the coefficient while ``ConstantSplit`` + scales as the second. This class raises the interpolated coefficient to ``power`` before the builder + squares it, so the effective exponent is ``2 * power`` and the default ``power=1.0`` matches + ``ConstantSplit``. Under the shared ``LogUniform(1e-6, 1e6)`` prior the prior now spans ``lambda^2``, and + the regularization matrix stays positive-definite to ``c ~ 1e6`` rather than collapsing from ``c ~ 1e4``. + + The **zeroth leg is unaffected**: ``BrightnessZeroth`` already squares its ``zeroth_coefficient`` exactly + once, so ``zeroth_coefficient`` keeps its meaning between the two classes. + + **Migration from ``AdaptSplitZeroth``.** The split leg's coefficient scale is squared: + ``c_new = c_old ** 2``; ``zeroth_coefficient`` is unchanged. To reproduce the legacy class exactly, pass + ``power=2.0``. + + **JAX & gradient support**: as for ``AdaptSplitZeroth`` -- FD-certified on the Delaunay mesh family (KNN + meshes), structurally incompatible with the rectangular meshes. + + Parameters + ---------- + zeroth_coefficient + The regularization coefficient of the zeroth-order leg, which is squared exactly once and is + therefore unchanged from ``AdaptSplitZeroth``. + zeroth_signal_scale + A factor which controls how rapidly the zeroth-order regularization varies from high signal regions + to low signal regions. + inner_coefficient + The inner regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the inner (high signal) regions of a mesh's reconstruction. + outer_coefficient + The outer regularization coefficient which controls the degree of smoothing of the inversion + reconstruction in the outer (low signal) regions of a mesh's reconstruction. + signal_scale + A factor which controls how rapidly the smoothness of regularization varies from high signal regions + to low signal regions. + power + The exponent the interpolated coefficient is raised to before the matrix builder squares it, so the + coefficient enters the split leg of the regularization matrix at the power ``2 * power``. The default + ``1.0`` is the ``ConstantSplit`` convention; ``2.0`` is the legacy ``AdaptSplitZeroth`` convention. + This is a convention switch, not a model parameter -- the shipped prior config fixes it as a + ``Constant`` prior so a search never samples it. + """ + super().__init__( + zeroth_coefficient=zeroth_coefficient, + zeroth_signal_scale=zeroth_signal_scale, + inner_coefficient=inner_coefficient, + outer_coefficient=outer_coefficient, + signal_scale=signal_scale, + ) + + self.power = power + + def regularization_weights_from(self, linear_obj: LinearObj, xp=np) -> np.ndarray: + """ + Returns the regularization weights of this regularization scheme. + + These are the interpolated inner / outer coefficients raised to ``self.power`` (default ``1.0``), as + opposed to ``AdaptSplitZeroth``, which squares them. + + Parameters + ---------- + linear_obj + The linear object (e.g. a ``Mapper``) which uses these weights when performing regularization. + + Returns + ------- + The regularization weights. + """ + pixel_signals = linear_obj.pixel_signals_from( + signal_scale=self.signal_scale, xp=xp + ) + + return adapt_regularization_weights_from( + inner_coefficient=self.inner_coefficient, + outer_coefficient=self.outer_coefficient, + pixel_signals=pixel_signals, + power=self.power, + ) diff --git a/autoarray/inversion/regularization/matern_adapt_kernel.py b/autoarray/inversion/regularization/matern_adapt_kernel.py index b00fb72d3..8db0ce6d4 100644 --- a/autoarray/inversion/regularization/matern_adapt_kernel.py +++ b/autoarray/inversion/regularization/matern_adapt_kernel.py @@ -52,7 +52,26 @@ def __init__( **JAX & gradient support**: as for ``MaternKernel`` (tfp ``bessel_kve`` gradients; explicit-inverse conditioning caveat). Note the defaults ``inner_coefficient == outer_coefficient == 1.0`` make - the weighting uniform — numerically identical to ``MaternKernel``. + the weighting uniform — but *not* numerically identical to ``MaternKernel``; see the + coefficient-convention note below. + + **Coefficient convention (legacy, ``lambda^4``).** The coefficients are squared twice before they + reach the regularization matrix -- once by ``adapt_regularization_weights_from`` and once by the + matrix builder -- so the matrix scales as the *fourth* power of the coefficient, while + ``Constant`` scales as the second. Both carry the same ``LogUniform(1e-6, 1e6)`` prior, so this + scheme explores a far wider effective smoothing range and reaches a numerically non + positive-definite matrix from ``c ~ 1e4`` where ``Constant`` survives to ``c ~ 1e6``. + + The adaptive weights enter the kernel covariance as + ``C_ij = K(d_ij) * w_i * w_j`` with ``w = 1 / regularization_weights``, so the regularization + matrix ``H = C^-1`` scales as the fourth power of the coefficient here too (``MaternKernel``, + by contrast, scales linearly in its ``coefficient``). + + This behaviour is preserved deliberately: changing it would alter the coefficient scale of every + adaptive fit ever run. **New work should prefer ``MaternAdaptPowerKernel``**, which takes a ``power`` argument + (default ``1.0``, giving the ``Constant``-matching ``lambda^2`` convention) and is more robust to + gradient / NaN pathologies. The migration is ``c_new = c_old ** 2``, and + ``MaternAdaptPowerKernel(power=2.0)`` reproduces this class's coefficient scaling exactly. Parameters ---------- diff --git a/autoarray/inversion/regularization/matern_adapt_power_kernel.py b/autoarray/inversion/regularization/matern_adapt_power_kernel.py new file mode 100644 index 000000000..a6cfee945 --- /dev/null +++ b/autoarray/inversion/regularization/matern_adapt_power_kernel.py @@ -0,0 +1,115 @@ +from __future__ import annotations +import numpy as np +from typing import Optional, TYPE_CHECKING + +if TYPE_CHECKING: + from autoarray.inversion.linear_obj.linear_obj import LinearObj + +from autoarray.inversion.regularization.adapt import adapt_regularization_weights_from +from autoarray.inversion.regularization.matern_adapt_kernel import MaternAdaptKernel + + +class MaternAdaptPowerKernel(MaternAdaptKernel): + def __init__( + self, + scale: float = 1.0, + nu: float = 0.5, + inner_coefficient: float = 1.0, + outer_coefficient: float = 1.0, + signal_scale: float = 1.0, + jitter: Optional[float] = None, + jitter_relative: bool = False, + power: float = 1.0, + ): + """ + Regularization which uses a Matern smoothing kernel with regularization weights that adapt to the + brightness of the source being reconstructed, with the coefficient convention of ``Constant``. + + This is the corrected sibling of ``MaternAdaptKernel``. The kernel, the covariance construction and the + jitter handling are all unchanged -- the only difference is the coefficient convention. + + ``MaternAdaptKernel`` squares its coefficients when interpolating them into per-pixel weights, and those + weights then enter the kernel covariance as ``C_ij = K(d_ij) * w_i * w_j`` with + ``w = 1 / regularization_weights``. The regularization matrix ``H = C^-1`` therefore scales as the fourth + power of the coefficient. This class raises the interpolated coefficient to ``power`` instead, so the + effective exponent is ``2 * power`` and the default ``power=1.0`` puts the coefficient into ``H`` at + ``lambda^2`` -- the ``Constant`` convention. Under the shared ``LogUniform(1e-6, 1e6)`` prior the prior + now spans ``lambda^2``. + + Note that ``MaternKernel`` itself scales *linearly* in its ``coefficient`` (``H = coefficient * C^-1``), + so neither this class nor ``MaternAdaptKernel`` reduces to it; the reference convention here is + ``Constant`` / ``AdaptPower``. + + **Migration from ``MaternAdaptKernel``.** The coefficient scale is squared: ``c_new = c_old ** 2``. To + reproduce the legacy class exactly, pass ``power=2.0``. + + **JAX & gradient support**: as for ``MaternAdaptKernel`` (tfp ``bessel_kve`` gradients; explicit-inverse + conditioning caveat). + + Parameters + ---------- + scale + The typical scale (correlation length) of the Matern regularization kernel. + nu + Controls the smoothness (differentiability) of the Matern kernel; ``nu=0.5`` corresponds to an + exponential (Ornstein-Uhlenbeck) kernel, while a Gaussian covariance is obtained in the limit as + ``nu`` approaches infinity. + inner_coefficient + The inner regularization coefficient which controls the degree of smoothing in the inner (high + signal) regions of a mesh's reconstruction. + outer_coefficient + The outer regularization coefficient which controls the degree of smoothing in the outer (low + signal) regions of a mesh's reconstruction. + signal_scale + A factor which controls how rapidly the smoothness of regularization varies from high signal regions + to low signal regions. + jitter + The small value added to the covariance diagonal for numerical stability. ``None`` (default) uses the + historical value 1e-8. + jitter_relative + If ``True`` the jitter is applied *relative* to each pixel's own variance (``C_ii *= 1 + jitter``) + rather than as a fixed absolute ``jitter * I``. See :func:`apply_jitter`. + power + The exponent the interpolated coefficient is raised to, so the coefficient enters the regularization + matrix at the power ``2 * power``. The default ``1.0`` is the ``Constant`` convention; ``2.0`` is the + legacy ``MaternAdaptKernel`` convention. This is a convention switch, not a model parameter -- the + shipped prior config fixes it as a ``Constant`` prior so a search never samples it. + """ + super().__init__( + scale=scale, + nu=nu, + inner_coefficient=inner_coefficient, + outer_coefficient=outer_coefficient, + signal_scale=signal_scale, + jitter=jitter, + jitter_relative=jitter_relative, + ) + + self.power = power + + def regularization_weights_from(self, linear_obj: LinearObj, xp=np) -> np.ndarray: + """ + Returns the regularization weights of this regularization scheme. + + These are the interpolated inner / outer coefficients raised to ``self.power`` (default ``1.0``), as + opposed to ``MaternAdaptKernel``, which squares them. + + Parameters + ---------- + linear_obj + The linear object (e.g. a ``Mapper``) which uses these weights when performing regularization. + + Returns + ------- + The regularization weights. + """ + pixel_signals = linear_obj.pixel_signals_from( + signal_scale=self.signal_scale, xp=xp + ) + + return adapt_regularization_weights_from( + inner_coefficient=self.inner_coefficient, + outer_coefficient=self.outer_coefficient, + pixel_signals=pixel_signals, + power=self.power, + ) diff --git a/autoarray/inversion/regularization/regularization_util.py b/autoarray/inversion/regularization/regularization_util.py index b3e32c9de..e38197c0b 100644 --- a/autoarray/inversion/regularization/regularization_util.py +++ b/autoarray/inversion/regularization/regularization_util.py @@ -9,6 +9,9 @@ from autoarray.inversion.regularization.adapt import ( weighted_regularization_matrix_from, ) +from autoarray.inversion.regularization.adapt import ( + weighted_regularization_matrix_single_scatter_from, +) from autoarray.inversion.regularization.brightness_zeroth import ( brightness_zeroth_regularization_matrix_from, ) diff --git a/test_autoarray/inversion/regularizations/test_adapt_power.py b/test_autoarray/inversion/regularizations/test_adapt_power.py new file mode 100644 index 000000000..2c2ed4f3e --- /dev/null +++ b/test_autoarray/inversion/regularizations/test_adapt_power.py @@ -0,0 +1,211 @@ +""" +The ``AdaptPower`` family: the corrected siblings of the ``Adapt`` family. + +Two things separate them from the legacy classes, and both are asserted here: + +1. the coefficient enters the regularization matrix at ``2 * power`` (default ``lambda^2``, the + ``Constant`` convention) rather than always at ``lambda^4``; +2. every mesh edge is scattered once rather than twice, so the non-split class equals ``Constant`` + exactly instead of being ``2 x`` it. + +The legacy classes are untouched — their own tests stay as they are. +""" + +import numpy as np +import pytest + +import autoarray as aa + + +@pytest.fixture(name="rectangular_mapper_9") +def make_rectangular_mapper_9(): + source_plane_mesh_grid = aa.Grid2D.no_mask( + values=[ + [0.1, 0.1], + [0.1, 0.2], + [0.1, 0.3], + [0.2, 0.1], + [0.2, 0.2], + [0.2, 0.3], + [0.3, 0.1], + [0.3, 0.2], + [0.3, 0.3], + ], + shape_native=(3, 3), + pixel_scales=1.0, + ) + + mesh_geometry = aa.MeshGeometryRectangular( + mesh=aa.mesh.RectangularUniform(shape=(3, 3)), + mesh_grid=source_plane_mesh_grid, + data_grid=None, + ) + + return aa.m.MockMapper( + source_plane_mesh_grid=source_plane_mesh_grid, + pixel_signals=np.array([1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]), + mesh_geometry=mesh_geometry, + ) + + +@pytest.mark.parametrize("coefficient", [0.5, 2.0, 7.0]) +def test__uniform_coefficients__matrix_equals_constant_exactly( + rectangular_mapper_9, coefficient +): + """ + The contract ``Adapt``'s docstring always claimed and never delivered: with + ``inner_coefficient == outer_coefficient`` the adaptive weighting is uniform, so the scheme must + reduce to ``Constant`` of the same coefficient. + """ + regularization_matrix = aa.reg.AdaptPower( + inner_coefficient=coefficient, outer_coefficient=coefficient + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + regularization_matrix_constant = aa.reg.Constant( + coefficient=coefficient + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + assert regularization_matrix == pytest.approx( + regularization_matrix_constant, abs=1.0e-12 + ) + + +def test__legacy_adapt_is_twice_constant__documents_the_scatter_asymmetry( + rectangular_mapper_9, +): + """ + A characterisation test for the legacy behaviour the ``*Power`` classes fix: ``Adapt`` scatters + every edge twice, so it is ``2 x`` ``Constant`` (up to the shared ``1e-8`` diagonal floor). + """ + floor = 1.0e-8 * np.eye(9) + + regularization_matrix_adapt = aa.reg.Adapt( + inner_coefficient=3.0, outer_coefficient=3.0 + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + regularization_matrix_constant = aa.reg.Constant( + coefficient=9.0 + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + assert regularization_matrix_adapt - floor == pytest.approx( + 2.0 * (regularization_matrix_constant - floor), rel=1.0e-12 + ) + + +def test__weights_are_unsquared_by_default__and_squared_at_power_2( + rectangular_mapper_9, +): + weights_legacy = aa.reg.Adapt( + inner_coefficient=1.0, outer_coefficient=2.0 + ).regularization_weights_from(linear_obj=rectangular_mapper_9) + + weights_power_1 = aa.reg.AdaptPower( + inner_coefficient=1.0, outer_coefficient=2.0 + ).regularization_weights_from(linear_obj=rectangular_mapper_9) + + weights_power_2 = aa.reg.AdaptPower( + inner_coefficient=1.0, outer_coefficient=2.0, power=2.0 + ).regularization_weights_from(linear_obj=rectangular_mapper_9) + + assert weights_power_1**2.0 == pytest.approx(weights_legacy, 1.0e-12) + assert (weights_power_2 == weights_legacy).all() + + +def test__power_2_reproduces_legacy_matrix_up_to_the_scatter_factor( + rectangular_mapper_9, +): + """ + ``power=2.0`` restores the legacy ``lambda^4`` coefficient scaling. The factor-2 scatter is fixed + regardless (that is the point of the new class), so the two matrices differ by exactly 2. + """ + floor = 1.0e-8 * np.eye(9) + + regularization_matrix_legacy = aa.reg.Adapt( + inner_coefficient=1.0, outer_coefficient=2.0 + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + regularization_matrix_power = aa.reg.AdaptPower( + inner_coefficient=1.0, outer_coefficient=2.0, power=2.0 + ).regularization_matrix_from(linear_obj=rectangular_mapper_9) + + assert regularization_matrix_legacy - floor == pytest.approx( + 2.0 * (regularization_matrix_power - floor), rel=1.0e-12 + ) + + +def test__migration__power_class_coefficient_is_the_legacy_coefficient_squared( + rectangular_mapper_9, +): + """ + The documented migration ``c_new = c_old ** 2``, checked on the weights. It is exact when the + inner and outer coefficients are equal; with differing coefficients the squaring happens after + the interpolation, so the two are not related term by term. + """ + weights_legacy_uniform = aa.reg.Adapt( + inner_coefficient=2.0, outer_coefficient=2.0 + ).regularization_weights_from(linear_obj=rectangular_mapper_9) + + weights_power_uniform = aa.reg.AdaptPower( + inner_coefficient=4.0, outer_coefficient=4.0 + ).regularization_weights_from(linear_obj=rectangular_mapper_9) + + assert weights_legacy_uniform == pytest.approx(weights_power_uniform, 1.0e-12) + + +def test__single_scatter_matrix__is_symmetric_positive_semi_definite(): + """ + With adaptive (non-uniform) weights the builder must still return a symmetric, positive + semi-definite matrix — it is a weighted graph Laplacian plus a ``1e-8`` diagonal floor, so every + row sums to the floor. + """ + neighbors = np.array( + [ + [1, 3, -1, -1], + [0, 2, 4, -1], + [1, 5, -1, -1], + [0, 4, 6, -1], + [1, 3, 5, 7], + [2, 4, 8, -1], + [3, 7, -1, -1], + [4, 6, 8, -1], + [5, 7, -1, -1], + ] + ) + + weights = np.array([0.1, 5.0, 0.3, 2.0, 9.0, 0.05, 1.0, 4.0, 0.7]) + + regularization_matrix = ( + aa.util.regularization.weighted_regularization_matrix_single_scatter_from( + regularization_weights=weights, neighbors=neighbors + ) + ) + + assert regularization_matrix == pytest.approx(regularization_matrix.T, 1.0e-12) + assert regularization_matrix.sum(axis=1) == pytest.approx( + 1.0e-8 * np.ones(9), abs=1.0e-14 + ) + assert np.linalg.eigvalsh(regularization_matrix).min() > 0.0 + + +def test__single_scatter_matrix__ignores_padded_neighbor_entries(): + """ + ``-1`` entries are padding. The symmetric edge weight is the mean of the two endpoints, so a + padded entry would contribute half of its own pixel's weight unless it is masked out. + """ + neighbors_padded = np.array([[1, -1], [0, -1]]) + neighbors_unpadded = np.array([[1], [0]]) + + weights = np.array([1.0, 4.0]) + + matrix_padded = ( + aa.util.regularization.weighted_regularization_matrix_single_scatter_from( + regularization_weights=weights, neighbors=neighbors_padded + ) + ) + matrix_unpadded = ( + aa.util.regularization.weighted_regularization_matrix_single_scatter_from( + regularization_weights=weights, neighbors=neighbors_unpadded + ) + ) + + assert matrix_padded == pytest.approx(matrix_unpadded, 1.0e-12) diff --git a/test_autoarray/inversion/regularizations/test_adapt_power_jax.py b/test_autoarray/inversion/regularizations/test_adapt_power_jax.py new file mode 100644 index 000000000..18604f532 --- /dev/null +++ b/test_autoarray/inversion/regularizations/test_adapt_power_jax.py @@ -0,0 +1,126 @@ +""" +JAX leg of the ``*Power`` gate: the new single-scatter matrix builder and the split path must agree +with their NumPy counterparts under ``jax.numpy``. + +Skipped when JAX is absent (it is an optional dependency). +""" + +import numpy as np +import pytest + +jax = pytest.importorskip("jax") +jnp = pytest.importorskip("jax.numpy") + +jax.config.update("jax_enable_x64", True) + +from autoarray.inversion.regularization.adapt import ( # noqa: E402 + weighted_regularization_matrix_single_scatter_from, +) +from autoarray.inversion.regularization.constant import ( # noqa: E402 + constant_regularization_matrix_from, +) +from autoarray.inversion.regularization import regularization_util # noqa: E402 + + +NEIGHBORS = np.array( + [ + [1, 3, -1, -1], + [0, 2, 4, -1], + [1, 5, -1, -1], + [0, 4, 6, -1], + [1, 3, 5, 7], + [2, 4, 8, -1], + [3, 7, -1, -1], + [4, 6, 8, -1], + [5, 7, -1, -1], + ] +) + + +def test__single_scatter_builder__numpy_and_jax_agree(): + weights = np.array([0.1, 5.0, 0.3, 2.0, 9.0, 0.05, 1.0, 4.0, 0.7]) + + matrix_np = weighted_regularization_matrix_single_scatter_from( + regularization_weights=weights, neighbors=NEIGHBORS + ) + matrix_jax = weighted_regularization_matrix_single_scatter_from( + regularization_weights=jnp.asarray(weights), + neighbors=jnp.asarray(NEIGHBORS), + xp=jnp, + ) + + assert np.asarray(matrix_jax) == pytest.approx(matrix_np, abs=1.0e-12) + + +def test__single_scatter_builder__jax_path_equals_constant_for_uniform_weights(): + coefficient = 3.0 + + matrix_jax = weighted_regularization_matrix_single_scatter_from( + regularization_weights=jnp.full(9, coefficient), + neighbors=jnp.asarray(NEIGHBORS), + xp=jnp, + ) + matrix_constant = constant_regularization_matrix_from( + coefficient=coefficient, + neighbors=NEIGHBORS.copy(), + neighbors_sizes=(NEIGHBORS >= 0).sum(axis=1), + ) + + assert np.asarray(matrix_jax) == pytest.approx(matrix_constant, abs=1.0e-12) + + +def test__single_scatter_builder__is_jittable_and_differentiable(): + def log_det(weights): + matrix = weighted_regularization_matrix_single_scatter_from( + regularization_weights=weights, + neighbors=jnp.asarray(NEIGHBORS), + xp=jnp, + ) + return jnp.linalg.slogdet(matrix)[1] + + weights = jnp.array([0.5, 1.5, 0.8, 2.0, 3.0, 0.4, 1.0, 1.2, 0.9]) + + value = jax.jit(log_det)(weights) + gradient = jax.grad(log_det)(weights) + + assert np.isfinite(np.asarray(value)) + assert np.all(np.isfinite(np.asarray(gradient))) + + +def test__split_builder__numpy_and_jax_agree(delaunay_mapper_9_3x3): + mappings, sizes, weights = ( + delaunay_mapper_9_3x3.interpolator._mappings_sizes_weights_split + ) + + regularization_weights = np.full(9, 2.0) + + def matrix_from(xp, mappings, sizes, weights, regularization_weights): + ( + splitted_mappings, + splitted_sizes, + splitted_weights, + ) = regularization_util.reg_split_from( + splitted_mappings=mappings, + splitted_sizes=sizes, + splitted_weights=weights, + xp=xp, + ) + + return regularization_util.pixel_splitted_regularization_matrix_from( + regularization_weights=regularization_weights, + splitted_mappings=splitted_mappings, + splitted_sizes=splitted_sizes, + splitted_weights=splitted_weights, + xp=xp, + ) + + matrix_np = matrix_from(np, mappings, sizes, weights, regularization_weights) + matrix_jax = matrix_from( + jnp, + jnp.asarray(mappings), + jnp.asarray(sizes), + jnp.asarray(weights), + jnp.asarray(regularization_weights), + ) + + assert np.asarray(matrix_jax) == pytest.approx(matrix_np, abs=1.0e-10) diff --git a/test_autoarray/inversion/regularizations/test_adapt_split_power.py b/test_autoarray/inversion/regularizations/test_adapt_split_power.py new file mode 100644 index 000000000..ba8d57338 --- /dev/null +++ b/test_autoarray/inversion/regularizations/test_adapt_split_power.py @@ -0,0 +1,138 @@ +""" +``AdaptSplitPower`` / ``AdaptSplitZerothPower``: the corrected siblings of the split adaptive schemes. + +The split family shares ``pixel_splitted_regularization_matrix_from`` with ``ConstantSplit``, so it +never carried ``Adapt``'s factor-2 scatter asymmetry. The only difference is the coefficient +convention, and with ``power=1.0`` (the default) the two families coincide exactly. +""" + +import numpy as np +import pytest + +import autoarray as aa + + +@pytest.mark.parametrize("coefficient", [0.5, 2.0, 7.0]) +def test__uniform_coefficients__matrix_equals_constant_split_exactly( + delaunay_mapper_9_3x3, coefficient +): + regularization_matrix = aa.reg.AdaptSplitPower( + inner_coefficient=coefficient, outer_coefficient=coefficient + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + regularization_matrix_constant_split = aa.reg.ConstantSplit( + coefficient=coefficient + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + assert (regularization_matrix == regularization_matrix_constant_split).all() + + +def test__power_2_reproduces_the_legacy_adapt_split_matrix(delaunay_mapper_9_3x3): + regularization_matrix_legacy = aa.reg.AdaptSplit( + inner_coefficient=1.0, outer_coefficient=2.0, signal_scale=1.0 + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + regularization_matrix_power = aa.reg.AdaptSplitPower( + inner_coefficient=1.0, outer_coefficient=2.0, signal_scale=1.0, power=2.0 + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + assert (regularization_matrix_power == regularization_matrix_legacy).all() + + +def test__weights_are_unsquared_by_default(delaunay_mapper_9_3x3): + weights_legacy = aa.reg.AdaptSplit( + inner_coefficient=1.0, outer_coefficient=2.0 + ).regularization_weights_from(linear_obj=delaunay_mapper_9_3x3) + + weights_power = aa.reg.AdaptSplitPower( + inner_coefficient=1.0, outer_coefficient=2.0 + ).regularization_weights_from(linear_obj=delaunay_mapper_9_3x3) + + assert weights_power**2.0 == pytest.approx(weights_legacy, 1.0e-12) + + +def test__zeroth__power_2_reproduces_the_legacy_matrix(delaunay_mapper_9_3x3): + regularization_matrix_legacy = aa.reg.AdaptSplitZeroth( + inner_coefficient=1.0, + outer_coefficient=2.0, + signal_scale=1.0, + zeroth_coefficient=3.0, + zeroth_signal_scale=2.0, + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + regularization_matrix_power = aa.reg.AdaptSplitZerothPower( + inner_coefficient=1.0, + outer_coefficient=2.0, + signal_scale=1.0, + zeroth_coefficient=3.0, + zeroth_signal_scale=2.0, + power=2.0, + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + assert (regularization_matrix_power == regularization_matrix_legacy).all() + + +def test__zeroth__split_leg_equals_constant_split_plus_the_unchanged_zeroth_leg( + delaunay_mapper_9_3x3, +): + """ + ``zeroth_coefficient`` is squared exactly once by ``BrightnessZeroth`` and is therefore unchanged + between the two classes — only the split leg's convention moves. + """ + regularization_matrix = aa.reg.AdaptSplitZerothPower( + inner_coefficient=2.0, + outer_coefficient=2.0, + signal_scale=1.0, + zeroth_coefficient=3.0, + zeroth_signal_scale=2.0, + ).regularization_matrix_from(linear_obj=delaunay_mapper_9_3x3) + + regularization_matrix_expected = aa.reg.ConstantSplit( + coefficient=2.0 + ).regularization_matrix_from( + linear_obj=delaunay_mapper_9_3x3 + ) + aa.reg.BrightnessZeroth( + coefficient=3.0, signal_scale=2.0 + ).regularization_matrix_from( + linear_obj=delaunay_mapper_9_3x3 + ) + + assert regularization_matrix == pytest.approx( + regularization_matrix_expected, 1.0e-12 + ) + + +def test__matern__weights_are_unsquared_by_default_and_squared_at_power_2(): + source_plane_mesh_grid = aa.Grid2D.no_mask( + values=[[0.1, 0.1], [1.1, 0.6], [2.1, 0.1], [0.4, 1.1], [1.1, 7.1], [2.1, 1.1]], + shape_native=(3, 2), + pixel_scales=1.0, + ) + + mapper = aa.m.MockMapper( + source_plane_mesh_grid=source_plane_mesh_grid, + pixel_signals=np.array([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]), + ) + + kwargs = dict( + scale=0.1, + nu=0.5, + inner_coefficient=0.1, + outer_coefficient=0.2, + signal_scale=0.1, + ) + + weights_legacy = aa.reg.MaternAdaptKernel(**kwargs).regularization_weights_from( + linear_obj=mapper + ) + + weights_power_1 = aa.reg.MaternAdaptPowerKernel( + **kwargs + ).regularization_weights_from(linear_obj=mapper) + + weights_power_2 = aa.reg.MaternAdaptPowerKernel( + power=2.0, **kwargs + ).regularization_weights_from(linear_obj=mapper) + + assert weights_power_1**2.0 == pytest.approx(weights_legacy, 1.0e-12) + assert (weights_power_2 == weights_legacy).all()