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
11 changes: 7 additions & 4 deletions autolens/point/solver/shape_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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 "
Expand Down
14 changes: 8 additions & 6 deletions test_autolens/point/triangles/test_shape_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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(
Expand All @@ -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."
Expand Down
Loading