Skip to content

Commit 5a6bd58

Browse files
authored
Merge pull request #378 from Jammy2211/feature/pure_callback_nfw
feature/pure_callback_nfw
2 parents f394907 + 038ec80 commit 5a6bd58

5 files changed

Lines changed: 20 additions & 8 deletions

File tree

autolens/analysis/analysis/dataset.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -85,6 +85,7 @@ def __init__(
8585
self=self,
8686
positions_likelihood_list=positions_likelihood_list,
8787
cosmology=cosmology,
88+
use_jax=use_jax
8889
)
8990

9091
self.raise_inversion_positions_likelihood_exception = (

autolens/analysis/analysis/lens.py

Lines changed: 15 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ def __init__(
2323
self,
2424
positions_likelihood_list: Optional[List[PositionsLH]] = None,
2525
cosmology: ag.cosmo.LensingCosmology = None,
26+
use_jax: bool = True,
2627
):
2728
"""
2829
Analysis classes are used by PyAutoFit to fit a model to a dataset via a non-linear search.
@@ -44,6 +45,15 @@ def __init__(
4445
self.cosmology = cosmology or Planck15()
4546
self.positions_likelihood_list = positions_likelihood_list
4647

48+
self._use_jax = use_jax
49+
50+
@property
51+
def _xp(self):
52+
if self._use_jax:
53+
import jax.numpy as jnp
54+
return jnp
55+
return np
56+
4757
def tracer_via_instance_from(
4858
self,
4959
instance: af.ModelInstance,
@@ -72,8 +82,9 @@ def tracer_via_instance_from(
7282
subhalo_centre = tracer_util.grid_2d_at_redshift_from(
7383
galaxies=instance.galaxies,
7484
redshift=instance.galaxies.subhalo.redshift,
75-
grid=aa.Grid2DIrregular(values=[instance.galaxies.subhalo.mass.centre]),
85+
grid=aa.Grid2DIrregular(values=[instance.galaxies.subhalo.mass.centre], xp=self._xp),
7686
cosmology=self.cosmology,
87+
xp=self._xp
7788
)
7889

7990
instance.galaxies.subhalo.mass.centre = tuple(subhalo_centre.in_list[0])
@@ -95,7 +106,7 @@ def tracer_via_instance_from(
95106
)
96107

97108
def log_likelihood_penalty_from(
98-
self, instance: af.ModelInstance, xp=np
109+
self, instance: af.ModelInstance,
99110
) -> Optional[float]:
100111
"""
101112
Call the positions overwrite log likelihood function, which add a penalty term to the likelihood if the
@@ -116,7 +127,7 @@ def log_likelihood_penalty_from(
116127
The penalty value of the positions log likelihood, if the positions do not trace close in the source plane,
117128
else a None is returned to indicate there is no penalty.
118129
"""
119-
log_likelihood_penalty = xp.array(0.0)
130+
log_likelihood_penalty = self._xp.array(0.0)
120131

121132
if self.positions_likelihood_list is not None:
122133

@@ -126,7 +137,7 @@ def log_likelihood_penalty_from(
126137

127138
log_likelihood_penalty = (
128139
positions_likelihood.log_likelihood_penalty_from(
129-
instance=instance, analysis=self, xp=xp
140+
instance=instance, analysis=self, xp=self._xp
130141
)
131142
)
132143

autolens/imaging/model/analysis.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -59,7 +59,6 @@ def log_likelihood_function(self, instance: af.ModelInstance) -> float:
5959

6060
log_likelihood_penalty = self.log_likelihood_penalty_from(
6161
instance=instance,
62-
xp=self._xp
6362
)
6463

6564
if self._use_jax:

autolens/interferometer/model/analysis.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -130,7 +130,7 @@ def log_likelihood_function(self, instance):
130130
"""
131131

132132
log_likelihood_penalty = self.log_likelihood_penalty_from(
133-
instance=instance, xp=self._xp
133+
instance=instance,
134134
)
135135

136136
return self.fit_from(instance=instance).figure_of_merit - log_likelihood_penalty

autolens/lens/tracer_util.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -179,6 +179,7 @@ def grid_2d_at_redshift_from(
179179
galaxies: List[ag.Galaxy],
180180
grid: aa.type.Grid2DLike,
181181
cosmology: ag.cosmo.LensingCosmology = None,
182+
xp=np,
182183
) -> aa.type.Grid2DLike:
183184
"""
184185
Returns a ray-traced grid of 2D Cartesian (y,x) coordinates, which accounts for multi-plane ray-tracing, at a
@@ -237,7 +238,7 @@ def grid_2d_at_redshift_from(
237238

238239
if plane_index_with_redshift:
239240
traced_grid_list = traced_grid_2d_list_from(
240-
planes=planes, grid=grid, cosmology=cosmology
241+
planes=planes, grid=grid, cosmology=cosmology, xp=xp
241242
)
242243

243244
return traced_grid_list[plane_index_with_redshift[0]]
@@ -249,7 +250,7 @@ def grid_2d_at_redshift_from(
249250
planes.insert(plane_index_insert, [ag.Galaxy(redshift=redshift)])
250251

251252
traced_grid_list = traced_grid_2d_list_from(
252-
planes=planes, grid=grid, cosmology=cosmology
253+
planes=planes, grid=grid, cosmology=cosmology, xp=xp
253254
)
254255

255256
return traced_grid_list[plane_index_insert]

0 commit comments

Comments
 (0)