Skip to content

Commit 7c975de

Browse files
authored
Merge pull request #384 from Jammy2211/feature/fft_jax_imaging
Feature/fft jax imaging
2 parents b05849c + f75db28 commit 7c975de

6 files changed

Lines changed: 14 additions & 22 deletions

File tree

autolens/__init__.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,7 @@
11
from autoconf import jax_wrapper
22
from autoconf.dictable import from_dict, from_json, output_to_json, to_dict
33
from autoarray import preprocess
4-
from autoarray.dataset.interferometer.w_tilde import (
5-
load_curvature_preload_if_compatible,
6-
)
7-
from autoarray.dataset.imaging.w_tilde import WTildeImaging
4+
85
from autoarray.dataset.imaging.dataset import Imaging
96
from autoarray.dataset.interferometer.dataset import (
107
Interferometer,

autolens/imaging/fit_imaging.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -112,7 +112,7 @@ def tracer_to_inversion(self) -> TracerToInversion:
112112
noise_map=self.noise_map,
113113
grids=self.grids,
114114
psf=self.dataset.psf,
115-
w_tilde=self.w_tilde,
115+
sparse_operator=self.dataset.sparse_operator,
116116
)
117117

118118
return TracerToInversion(

autolens/interferometer/fit_interferometer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,7 @@ def tracer_to_inversion(self) -> TracerToInversion:
115115
noise_map=self.noise_map,
116116
grids=self.grids,
117117
transformer=self.dataset.transformer,
118-
w_tilde=self.w_tilde,
118+
sparse_operator=self.dataset.sparse_operator,
119119
)
120120

121121
return TracerToInversion(

autolens/lens/to_inversion.py

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -175,18 +175,13 @@ def lp_linear_func_list_galaxy_dict(
175175
blurring=traced_blurring_grids_of_planes_list[plane_index],
176176
)
177177

178-
if self.dataset.w_tilde is not None:
179-
w_tilde = self.dataset.w_tilde
180-
else:
181-
w_tilde = None
182-
183178
dataset = aa.DatasetInterface(
184179
data=self.dataset.data,
185180
noise_map=self.dataset.noise_map,
186181
grids=grids,
187182
psf=self.psf,
188183
transformer=self.transformer,
189-
w_tilde=w_tilde,
184+
sparse_operator=self.dataset.sparse_operator,
190185
)
191186

192187
galaxies_to_inversion = ag.GalaxiesToInversion(

autolens/lens/tracer_util.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -172,7 +172,9 @@ def traced_grid_2d_list_from(
172172
)
173173

174174
# Remove NaN deflection values to sanitize the ray-tracing calculation for JAX.
175-
deflections_yx_2d = xp.where(xp.isfinite(deflections_yx_2d.array), deflections_yx_2d.array, 0.0)
175+
deflections_yx_2d = xp.where(
176+
xp.isfinite(deflections_yx_2d.array), deflections_yx_2d.array, 0.0
177+
)
176178

177179
traced_deflection_list.append(deflections_yx_2d)
178180

@@ -348,9 +350,7 @@ def time_delays_from(
348350
)
349351

350352
# Time-delay distance in meters: (1+z_l) * Dd * Ds / Dds
351-
D_dt_m = (
352-
(1.0 + z_l) * (Dd_kpc * Ds_kpc / Dds_kpc) * kpc_in_m
353-
)
353+
D_dt_m = (1.0 + z_l) * (Dd_kpc * Ds_kpc / Dds_kpc) * kpc_in_m
354354

355355
# Fermat potential (should be in arcsec^2 for this formula)
356356
fermat_potential = galaxies.fermat_potential_from(grid=grid, xp=xp)

test_autolens/imaging/test_simulate_and_fit_imaging.py

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -651,7 +651,7 @@ def test__simulate_imaging_data_and_fit__linear_light_profiles_and_pixelization_
651651
assert fit_linear.figure_of_merit == pytest.approx(-180.8284970580511, 1.0e-4)
652652

653653

654-
def test__simulate_imaging_data_and_fit__complex_fit_compare_mapping_matrix_w_tilde():
654+
def test__simulate_imaging_data_and_fit__complex_fit_compare_mapping_matrix_sparse_operator():
655655

656656
grid = al.Grid2D.uniform(shape_native=(21, 21), pixel_scales=0.1)
657657

@@ -735,20 +735,20 @@ def test__simulate_imaging_data_and_fit__complex_fit_compare_mapping_matrix_w_ti
735735
tracer=tracer,
736736
)
737737

738-
masked_dataset_w_tilde = masked_dataset.apply_w_tilde()
738+
masked_dataset_sparse_operator = masked_dataset.apply_sparse_operator()
739739

740-
fit_w_tilde = al.FitImaging(
741-
dataset=masked_dataset_w_tilde,
740+
fit_sparse_operator = al.FitImaging(
741+
dataset=masked_dataset_sparse_operator,
742742
tracer=tracer,
743743
)
744744

745745
assert fit_mapping.inversion.curvature_matrix == pytest.approx(
746-
fit_w_tilde.inversion.curvature_matrix,
746+
fit_sparse_operator.inversion.curvature_matrix,
747747
1.0e-4,
748748
)
749749

750750
assert fit_mapping.inversion.regularization_matrix == pytest.approx(
751-
fit_w_tilde.inversion.regularization_matrix,
751+
fit_sparse_operator.inversion.regularization_matrix,
752752
1.0e-4,
753753
)
754754

0 commit comments

Comments
 (0)