diff --git a/autolens/point/solver/point_solver.py b/autolens/point/solver/point_solver.py index a654fefa0..282f91e9a 100644 --- a/autolens/point/solver/point_solver.py +++ b/autolens/point/solver/point_solver.py @@ -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 ------- @@ -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 diff --git a/test_autolens/point/triangles/test_solver_edge_cases.py b/test_autolens/point/triangles/test_solver_edge_cases.py index 46f3972f8..85a2ac140 100644 --- a/test_autolens/point/triangles/test_solver_edge_cases.py +++ b/test_autolens/point/triangles/test_solver_edge_cases.py @@ -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 + ), ) @@ -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)