diff --git a/autoarray/structures/triangles/array.py b/autoarray/structures/triangles/array.py index 0ed7cd4e6..1e92897f8 100644 --- a/autoarray/structures/triangles/array.py +++ b/autoarray/structures/triangles/array.py @@ -10,7 +10,7 @@ from autoarray.structures.triangles.shape import Shape from autoarray.structures.triangles.shape import _barycentric_contains -MAX_CONTAINING_SIZE = 15 +MAX_CONTAINING_SIZE = 20 # Private A/B switch for the step-0 containment route (point-source CPU phase 4b, # PyAutoArray#579). It only affects `ArrayTriangles` built by diff --git a/test_autoarray/structures/triangles/test_coordinate_jax.py b/test_autoarray/structures/triangles/test_coordinate_jax.py index 292e831c7..2824d10aa 100644 --- a/test_autoarray/structures/triangles/test_coordinate_jax.py +++ b/test_autoarray/structures/triangles/test_coordinate_jax.py @@ -586,3 +586,61 @@ def containing(vertices, point): assert all(size < 3 * n * 2 for size in gathered), ( f"step-0 containment gathers {gathered} elements; the triangle array is {3 * n * 2}" ) + + +# A 16- and a 17-sheet fold: the static lattice's vertices are traced through z -> z**k (z = x + iy), +# so a source point at |w| = 1 has k preimages on the unit circle, each in the interior of exactly one +# traced triangle -- a k-member containing set with no boundary ties. The PointSolver's uncapped +# step-0 maximum on 200 prior draws was 17 (prior draw 12, point-source CPU phase 4c, +# PyAutoArray#583), which the old cap of 15 truncated silently. +FOLD_LIMITS = dict(y_min=-1.3, y_max=1.3, x_min=-1.3, x_max=1.3, scale=0.05) + + +@pytest.mark.parametrize("sheets", [16, 17]) +def test__containing_set_above_the_old_cap_is_not_truncated(sheets): + """ + Under ``jit``, a 16- or 17-triangle containing set keeps every member at the default + `MAX_CONTAINING_SIZE` (20), padded to ``(20,)``; the old cap of 15, passed explicitly, drops + some. Red at ``MAX_CONTAINING_SIZE = 15``. + """ + lattice = _static_lattice(FOLD_LIMITS) + z = lattice.vertices[:, 1] + 1j * lattice.vertices[:, 0] + w = z**sheets + traced = jnp.stack([w.imag, w.real], axis=1) + point = (float(np.sin(0.3)), float(np.cos(0.3))) + + uncapped = np.flatnonzero( + np.asarray(Point(*point).mask(lattice.with_vertices(traced).triangles)) + ) + assert len(uncapped) == sheets + + def containing(vertices, max_containing_size=None): + kwargs = ( + {} + if max_containing_size is None + else {"max_containing_size": max_containing_size} + ) + triangles = triangles_array.ArrayTriangles( + indices=lattice.indices, + vertices=vertices, + step0_layout=lattice.step0_layout, + **kwargs, + ) + return triangles.containing_indices(Point(*point)) + + assert lattice.with_vertices(traced).max_containing_size == MAX_CONTAINING_SIZE + + kept = np.asarray(jax.jit(containing)(traced)) + assert np.array_equal(np.sort(kept[kept >= 0]), uncapped) + assert kept.shape == (20,) + + kept_via_lattice = np.asarray( + jax.jit(lambda v: lattice.with_vertices(v).containing_indices(Point(*point)))( + traced + ) + ) + assert np.array_equal(kept_via_lattice, kept) + + old_cap = np.asarray(jax.jit(lambda v: containing(v, 15))(traced)) + assert old_cap.shape == (15,) + assert np.sum(old_cap >= 0) == 15 < sheets