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
Original file line number Diff line number Diff line change
Expand Up @@ -492,26 +492,6 @@ def curvature_matrix_with_added_to_diag_from(
return curvature_matrix


@numba_util.jit()
def curvature_matrix_mirrored_from(
curvature_matrix: np.ndarray,
) -> np.ndarray:
curvature_matrix_mirrored = np.zeros(
(curvature_matrix.shape[0], curvature_matrix.shape[1])
)

for i in range(curvature_matrix.shape[0]):
for j in range(curvature_matrix.shape[1]):
if curvature_matrix[i, j] != 0:
curvature_matrix_mirrored[i, j] = curvature_matrix[i, j]
curvature_matrix_mirrored[j, i] = curvature_matrix[i, j]
if curvature_matrix[j, i] != 0:
curvature_matrix_mirrored[i, j] = curvature_matrix[j, i]
curvature_matrix_mirrored[j, i] = curvature_matrix[j, i]

return curvature_matrix_mirrored


@numba_util.jit()
def curvature_matrix_via_sparse_operator_from(
psf_precision_operator: np.ndarray,
Expand Down Expand Up @@ -694,6 +674,72 @@ def curvature_matrix_off_diags_via_sparse_operator_from(
return curvature_matrix


@numba_util.jit()
def curvature_matrix_off_diags_via_mapper_and_blurred_curvature_weights_from(
data_to_pix_unique: np.ndarray,
data_weights: np.ndarray,
pix_lengths: np.ndarray,
pix_pixels: int,
blurred_curvature_weights: np.ndarray, # shape (n_unmasked, n_funcs)
) -> np.ndarray:
"""
Returns the off-diagonal terms in the curvature matrix `F` (see Warren & Dye 2003)
between a mapper object and a linear func object, from curvature weights that have
already been correlated with the PSF.

This is the scatter half of
`curvature_matrix_off_diags_via_mapper_and_linear_func_curvature_vector_from`: that
function expands the curvature weights onto the native grid, performs a dense
sliding-window correlation with the PSF and then runs this loop. Splitting the two lets
the correlation be done once per linear func by the batched FFT convolver (which is over
an order of magnitude faster at HST resolution) while this loop, which is genuinely
sparse and irregular, stays in numba and runs once per (mapper, linear func) pair.

For each unique mapping between a data pixel and a pixelization pixel, the PSF-correlated
curvature weights at that data pixel are multiplied by the mapping weight and accumulated
into the off-diagonal block of the curvature matrix. This accounts for sub-pixel mappings
between data pixels and pixelization pixels.

Parameters
----------
data_to_pix_unique
An array that maps every data pixel index (e.g. the masked image pixel indexes in 1D)
to its unique set of pixelization pixel indexes (see `data_slim_to_pixelization_unique_from`).
data_weights
For every unique mapping between a set of data sub-pixels and a pixelization pixel,
the weight of this mapping based on the number of sub-pixels that map to the pixelization pixel.
pix_lengths
A 1D array describing how many unique pixels each data pixel maps to. Used to iterate over
`data_to_pix_unique` and `data_weights`.
pix_pixels
The total number of pixels in the pixelization that reconstructs the data.
blurred_curvature_weights
The operated values of the linear function divided by the noise-map squared and
correlated with the PSF, with shape [n_unmasked_data_pixels, n_linear_func_pixels].

Returns
-------
ndarray
The off-diagonal block of the curvature matrix `F` (see Warren & Dye 2003),
with shape [pix_pixels, n_linear_func_pixels].
"""
data_pixels = data_weights.shape[0]
n_funcs = blurred_curvature_weights.shape[1]

off_diag = np.zeros((pix_pixels, n_funcs))

for data_0 in range(data_pixels):
for pix_0_index in range(pix_lengths[data_0]):
data_0_weight = data_weights[data_0, pix_0_index]
pix_0 = data_to_pix_unique[data_0, pix_0_index]
for f in range(n_funcs):
off_diag[pix_0, f] += (
data_0_weight * blurred_curvature_weights[data_0, f]
)

return off_diag


@numba_util.jit()
def curvature_matrix_off_diags_via_mapper_and_linear_func_curvature_vector_from(
data_to_pix_unique: np.ndarray,
Expand All @@ -714,6 +760,13 @@ def curvature_matrix_off_diags_via_mapper_and_linear_func_curvature_vector_from(
noise-map squared) are expanded into the native 2D image grid, convolved with the PSF
kernel, and then remapped back to the 1D slim representation.

The inversion itself no longer calls this function: the correlation is performed by the
batched FFT convolver (`Convolver.reversed_kernel`) and only the scatter/accumulate loop
below runs in numba, via
`curvature_matrix_off_diags_via_mapper_and_blurred_curvature_weights_from`. This dense
sliding-window implementation is retained as the reference the FFT path is asserted
against in the unit tests.

For each unique mapping between a data pixel and a pixelization pixel, the convolved
curvature weights at that data pixel are multiplied by the mapping weights and
accumulated into the off-diagonal block of the curvature matrix. This accounts for
Expand Down
161 changes: 139 additions & 22 deletions autoarray/inversion/inversion/imaging_numba/sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,10 +332,14 @@ def curvature_matrix(self) -> np.ndarray:
for simultaneously. In the w-tilde formalism this requires us to consider the mappings between data and every
linear object, meaning that the linear alegbra has both on and off diagonal terms.

The `curvature_matrix` computed here is overwritten in memory when the regularization matrix is added to it,
because for large matrices this avoids overhead. For this reason, `curvature_matrix` is not a cached property
to ensure if we access it after computing the `curvature_reg_matrix` it is correctly recalculated in a new
array of memory.
Every block of F is written into the matrix already symmetrized: the mapper x mapper blocks are
folded and mirrored inside `curvature_matrix_via_sparse_operator_from`, and the off-diagonal
blocks (mapper x mapper, mapper x linear-func, linear-func x linear-func) are each placed
together with their transpose. A global symmetrizing pass over the assembled matrix would
therefore be a no-op, and is not run.

`curvature_matrix` is a cached property, and `curvature_reg_matrix` adds the regularization
matrix to it out-of-place, so the cached F is never overwritten by that addition.
"""
if self.has(cls=AbstractLinearObjFuncList):
curvature_matrix = self._curvature_matrix_func_list_and_mapper
Expand All @@ -344,10 +348,6 @@ def curvature_matrix(self) -> np.ndarray:
else:
curvature_matrix = self._curvature_matrix_multi_mapper

curvature_matrix = inversion_imaging_numba_util.curvature_matrix_mirrored_from(
curvature_matrix=curvature_matrix,
)

if len(self.no_regularization_index_list) > 0:
curvature_matrix = (
inversion_imaging_numba_util.curvature_matrix_with_added_to_diag_from(
Expand Down Expand Up @@ -386,11 +386,9 @@ def _curvature_matrix_mapper_diag(self) -> Optional[np.ndarray]:
psf_precision_operator=self.sparse_operator.psf_precision_operator_sparse,
psf_precision_indexes=self.sparse_operator.indexes,
psf_precision_lengths=self.sparse_operator.lengths,
data_to_pix_unique=np.array(
mapper_i.unique_mappings.data_to_pix_unique
),
data_weights=np.array(mapper_i.unique_mappings.data_weights),
pix_lengths=np.array(mapper_i.unique_mappings.pix_lengths),
data_to_pix_unique=mapper_i.unique_mappings.data_to_pix_unique,
data_weights=mapper_i.unique_mappings.data_weights,
pix_lengths=mapper_i.unique_mappings.pix_lengths,
pix_pixels=mapper_i.params,
)

Expand Down Expand Up @@ -493,6 +491,11 @@ def _curvature_matrix_multi_mapper(self) -> np.ndarray:
mapper_param_range_j[0] : mapper_param_range_j[1],
] = off_diag

curvature_matrix[
mapper_param_range_j[0] : mapper_param_range_j[1],
mapper_param_range_i[0] : mapper_param_range_i[1],
] = off_diag.T

return curvature_matrix

@property
Expand All @@ -505,45 +508,159 @@ def _curvature_matrix_func_list_and_mapper(self) -> np.ndarray:
curvature matrix given by equation (4) and the letter F.

This function computes the diagonal terms of F using the sparse_operator formalism.

The three blocks of F are assembled by separate private helpers, so that each block can be
computed (and therefore profiled) on its own:

- the mapper x mapper block, `_curvature_matrix_mapper_diag` (via `_curvature_matrix_multi_mapper`);
- the mapper x linear-func blocks, `_curvature_matrix_mapper_func_blocks_from`;
- the linear-func x linear-func blocks, `_curvature_matrix_func_func_blocks_from`.

The helpers write into the `curvature_matrix` they are passed and return it, so composing them
in this order is exactly the single-pass assembly they replaced.
"""

curvature_matrix = self._curvature_matrix_multi_mapper

curvature_matrix = self._curvature_matrix_mapper_func_blocks_from(
curvature_matrix=curvature_matrix
)

curvature_matrix = self._curvature_matrix_func_func_blocks_from(
curvature_matrix=curvature_matrix
)

return curvature_matrix

def _blurred_curvature_weights_from(
self, curvature_weights: np.ndarray
) -> np.ndarray:
"""
Returns a linear func's noise-weighted curvature weights correlated with the PSF, in the
mask's slim representation with shape [n_unmasked_data_pixels, n_linear_func_pixels].

The mapper x linear-func block of `F` requires, at every unmasked data pixel, the
sliding-window sum ``sum_dy_dx psf[dy, dx] * weights[y + dy - cy, x + dx - cx]`` -- a
*correlation* with the PSF, not a convolution. Correlating with the PSF is exactly
convolving with the PSF reversed along both axes, so this routes through the dataset
PSF's `reversed_kernel` convolver and its batched (multi-column) convolution, which is
over an order of magnitude faster at HST resolution than the dense sliding window it
replaces.

As in the sliding-window implementation the weights are zero everywhere outside the
mask (no blurring mapping matrix is supplied), and the result is read back only at the
unmasked pixels.

Parameters
----------
curvature_weights
The operated values of a linear function divided by the noise-map squared, with
shape [n_unmasked_data_pixels, n_linear_func_pixels].
"""
return self.psf.reversed_kernel.convolved_mapping_matrix_from(
mapping_matrix=curvature_weights,
mask=self.mask,
xp=np,
)

def _curvature_matrix_mapper_func_blocks_from(
self, curvature_matrix: np.ndarray
) -> np.ndarray:
"""
Writes the mapper x linear-func off-diagonal blocks of the `curvature_matrix` into the input
matrix, returning it.

Each block contracts a mapper's unique data-to-source-pixel mappings against the PSF-correlated,
noise-weighted curvature vector of a linear function.

The correlation is the dominant cost of F at HST resolution, so it is done once per linear
func by `_blurred_curvature_weights_from` (batched FFT convolution) rather than once per
(mapper, linear func) pair by a dense sliding window, and only the sparse scatter of the
result onto source pixels runs in numba.

Each `[mapper, linear_func]` block is written together with its transpose into the
`[linear_func, mapper]` block, so F leaves this helper symmetric and no global mirroring
pass is required.

Parameters
----------
curvature_matrix
The (total_params, total_params) curvature matrix the blocks are written into.
"""

mapper_list = self.cls_list_from(cls=Mapper)
mapper_param_range_list = self.param_range_list_from(cls=Mapper)

if len(mapper_list) == 0:
return curvature_matrix

linear_func_list = self.cls_list_from(cls=AbstractLinearObjFuncList)
linear_func_param_range_list = self.param_range_list_from(
cls=AbstractLinearObjFuncList
)

# Neither the noise-weighted curvature weights of a linear func nor their PSF
# correlation depend on the mapper, so both are formed once per linear func rather
# than once per (mapper, linear func) pair.
blurred_curvature_weights_list = [
self._blurred_curvature_weights_from(
curvature_weights=np.array(
self.linear_func_operated_mapping_matrix_dict[linear_func]
/ self.noise_map[:, None] ** 2
)
)
for linear_func in linear_func_list
]

for i in range(len(mapper_list)):
mapper = mapper_list[i]
mapper_param_range = mapper_param_range_list[i]

for func_index, linear_func in enumerate(linear_func_list):
linear_func_param_range = linear_func_param_range_list[func_index]

data_linear_func_matrix = (
self.linear_func_operated_mapping_matrix_dict[linear_func]
/ self.noise_map[:, None] ** 2
)

off_diag = inversion_imaging_numba_util.curvature_matrix_off_diags_via_mapper_and_linear_func_curvature_vector_from(
off_diag = inversion_imaging_numba_util.curvature_matrix_off_diags_via_mapper_and_blurred_curvature_weights_from(
data_to_pix_unique=mapper.unique_mappings.data_to_pix_unique,
data_weights=mapper.unique_mappings.data_weights,
pix_lengths=mapper.unique_mappings.pix_lengths,
pix_pixels=mapper.params,
curvature_weights=np.array(data_linear_func_matrix),
mask=self.mask.array,
psf_kernel=self.psf.kernel.native.array,
blurred_curvature_weights=blurred_curvature_weights_list[
func_index
],
)

curvature_matrix[
mapper_param_range[0] : mapper_param_range[1],
linear_func_param_range[0] : linear_func_param_range[1],
] = off_diag

curvature_matrix[
linear_func_param_range[0] : linear_func_param_range[1],
mapper_param_range[0] : mapper_param_range[1],
] = off_diag.T

return curvature_matrix

def _curvature_matrix_func_func_blocks_from(
self, curvature_matrix: np.ndarray
) -> np.ndarray:
"""
Writes the linear-func x linear-func blocks of the `curvature_matrix` into the input matrix,
returning it.

Each block is a BLAS `dot` of two noise-weighted operated mapping matrices.

Parameters
----------
curvature_matrix
The (total_params, total_params) curvature matrix the blocks are written into.
"""

linear_func_list = self.cls_list_from(cls=AbstractLinearObjFuncList)
linear_func_param_range_list = self.param_range_list_from(
cls=AbstractLinearObjFuncList
)

# The linear func x linear func block is symmetric, so each weighted matrix is
# formed once and only the upper triangle of blocks is computed, with the
# mirrored block set from the transpose.
Expand Down
Loading
Loading