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
6 changes: 4 additions & 2 deletions autolens/point/solver/point_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,7 +82,9 @@ def solve(
Whether to strip the ``inf`` sentinel rows from the output. When ``None`` (the default),
defaults to ``True`` on the NumPy path and ``False`` on the JAX path. The JAX path
keeps the padded static shape so the output crosses a ``jax.jit`` boundary cleanly;
strip the infinities outside the jit if needed.
strip the infinities outside the jit if needed. The default follows the effective
``xp`` for this call, including an explicit override of the constructor backend.
An explicit ``True`` or ``False`` takes precedence over this default.

Returns
-------
Expand Down Expand Up @@ -119,7 +121,7 @@ def solve(
xp = self._xp

if remove_infinities is None:
remove_infinities = not self.use_jax
remove_infinities = xp is np

# NOTE: pytree registration is the user's responsibility (call
# `autolens.jax.register_tracer_classes(tracer)` once before wrapping
Expand Down
36 changes: 35 additions & 1 deletion test_autolens/point/triangles/test_solver_edge_cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,9 @@ def solver_grid():
def lens_galaxy():
return al.Galaxy(
redshift=0.5,
mass=al.mp.Isothermal(centre=(0.0, 0.0), ell_comps=(0.1, 0.0), einstein_radius=1.0),
mass=al.mp.Isothermal(
centre=(0.0, 0.0), ell_comps=(0.1, 0.0), einstein_radius=1.0
),
)


Expand Down Expand Up @@ -166,3 +168,35 @@ def test__precision_fine_enough__still_solves(solver_grid, lens_galaxy):
)

assert len(result) == 4


@pytest.mark.parametrize("use_jax", [False, True])
@pytest.mark.parametrize("remove_infinities", [None, False, True])
def test__numpy_override_controls_default_padding(
solver_grid, lens_galaxy, use_jax, remove_infinities
):
"""An explicit NumPy call strips rejected rows unless padding is requested.

use_jax=True only sets the constructor preference: this test executes NumPy
throughout. The opposite override and JIT are covered in the test workspace.
"""
source = (0.05, 0.02)
solver = al.PointSolver.for_grid(
grid=solver_grid,
pixel_scale_precision=0.001,
magnification_threshold=1.0e100,
use_jax=use_jax,
)
result = np.asarray(
solver.solve(
tracer=_tracer(lens_galaxy, source),
source_plane_coordinate=source,
xp=np,
remove_infinities=remove_infinities,
).array
)
if remove_infinities is False:
assert len(result) > 0
assert np.isinf(result).all()
else:
assert result.shape == (0, 2)
Loading