diff --git a/autolens/point/solver/shape_solver.py b/autolens/point/solver/shape_solver.py index 1f35b01cb..e5ac6b650 100644 --- a/autolens/point/solver/shape_solver.py +++ b/autolens/point/solver/shape_solver.py @@ -21,6 +21,7 @@ import autoarray as aa +from autoarray.structures.triangles.array import MAX_CONTAINING_SIZE from autoarray.structures.triangles.shape import Shape import autogalaxy as ag @@ -582,8 +583,9 @@ class ShapeSolver(AbstractSolver): --- ``ShapeSolver`` is a NumPy-only solver and rejects ``use_jax=True`` / ``xp=jax.numpy``. The JAX triangle containers keep static shapes by truncating every refinement step to - ``ArrayTriangles.MAX_CONTAINING_SIZE`` (15) triangles — ample for a ``Point``, which - lies inside a handful of triangles, and meaningless for a shape with area, whose kept + ``autoarray.structures.triangles.array.MAX_CONTAINING_SIZE`` triangles (the cap the + `NotImplementedError` message quotes) — ample for a ``Point``, which lies inside a + handful of triangles, and meaningless for a shape with area, whose kept set grows with the size of its images. Before this was found, ``use_jax=True`` was silently ignored (``find_magnification`` hardcoded ``xp=np``); routing it through ``self._xp`` instead would have replaced a silently-ignored flag with a silently wrong @@ -594,10 +596,11 @@ class ShapeSolver(AbstractSolver): # The one message both rejection routes raise, so a caller sees the same explanation # whether they set `use_jax=True` or passed `xp=jax.numpy` -- and, since it is a plain - # string, whether or not JAX is installed. + # string, whether or not JAX is installed. The cap is read from the (NumPy-only) + # `autoarray.structures.triangles.array` module, so the number cannot drift from it. _JAX_REJECTED_MESSAGE = ( "ShapeSolver does not support JAX. The JAX triangle containers truncate " - "every refinement step to ArrayTriangles.MAX_CONTAINING_SIZE (15) " + f"every refinement step to MAX_CONTAINING_SIZE ({MAX_CONTAINING_SIZE}) " "triangles to keep static shapes, which is enough for a Point but not for " "a Shape with area: the kept triangles of an extended source number in the " "thousands, so the JAX path silently measures a small fraction of the " diff --git a/test_autolens/point/triangles/test_shape_solver.py b/test_autolens/point/triangles/test_shape_solver.py index e90887568..0739bf4b6 100644 --- a/test_autolens/point/triangles/test_shape_solver.py +++ b/test_autolens/point/triangles/test_shape_solver.py @@ -20,6 +20,7 @@ import autoarray as aa import autolens as al +from autoarray.structures.triangles.array import MAX_CONTAINING_SIZE from autoarray.structures.triangles.shape import Circle, Polygon, Square, Triangle from autolens.point.solver.shape_solver import ShapeSolver @@ -575,7 +576,7 @@ def test_use_jax_solver_is_rejected_rather_than_silently_wrong(image_grid, sis_t with pytest.raises(NotImplementedError) as exc_info: call() - assert "MAX_CONTAINING_SIZE" in str(exc_info.value) + assert f"MAX_CONTAINING_SIZE ({MAX_CONTAINING_SIZE})" in str(exc_info.value) def test_explicit_jax_module_is_rejected_rather_than_silently_wrong( @@ -598,17 +599,18 @@ def test_explicit_jax_module_is_rejected_rather_than_silently_wrong( tracer=sis_tracer, shape=Circle(0.3, 0.0, radius=0.1), xp=jnp ) - assert "MAX_CONTAINING_SIZE" in str(exc_info.value) + assert f"MAX_CONTAINING_SIZE ({MAX_CONTAINING_SIZE})" in str(exc_info.value) @pytest.mark.xfail( strict=True, reason=( "DEFERRED: the JAX triangle containers truncate every refinement step to " - "ArrayTriangles.MAX_CONTAINING_SIZE (15) triangles to keep static shapes. That is " - "enough for a Point but not for a Shape with area, so the JAX path keeps 15 of the " - "~800 triangles the NumPy path keeps and measures a magnification of 0.13 where " - "the truth is 6.86. Lifting the cap is a redesign of the JAX containers (the cap " + f"MAX_CONTAINING_SIZE ({MAX_CONTAINING_SIZE}) triangles to keep static shapes. That " + "is enough for a Point but not for a Shape with area, so the JAX path keeps " + f"{MAX_CONTAINING_SIZE} of the " + "~800 triangles the NumPy path keeps and measures a magnification orders of " + "magnitude too small (0.13 at the original cap of 15, where the truth is 6.86). Lifting the cap is a redesign of the JAX containers (the cap " "is what makes their shapes static, and an extended source has no static bound), " "not a fix in ShapeSolver, so it is deferred and ShapeSolver raises on use_jax " "instead. Remove this xfail when the containers grow a dynamic kept set."