Skip to content

fix(inversion): zero-signal adapt image gives NaN on the JAX path #548

Description

@Jammy2211

Overview

adaptive_pixel_signals_from normalises the pixel signals by their maximum with
xp.where(max_sig > 0, pixel_signals / max_sig, pixel_signals). The where guards the
selection but the division is still evaluated for every element, so a zero-signal adapt
image — one whose emission lands nowhere near the pixels being adapted — makes max_sig
exactly 0 and produces a 0/0 NaN in the unselected branch. Found during the phase-6
rebuild of the ci-timing-fast-tests epic (autolens_workspace_test#293 / #294).

Investigating it sharpened the diagnosis: the two backends already agree in the forward
pass
— both where calls select the zeros, so a jnp forward result is finite. What
actually diverges is (a) NumPy's discarded RuntimeWarning: invalid value encountered in divide and (b) the gradient: under jax.grad the NaN propagates out of the discarded
branch and into the likelihood. That is the standard where-inside-grad trap, and it is
the path fitness._vmap takes before collapsing to resample_figure_of_merit.

Plan

  • Fix the normalisation to use a safe denominator rather than a where around the
    quotient, matching the idiom already used for pixel_counts two lines above.
  • Apply the same guard to the step-8 exponentiation, where 0.0 ** signal_scale is finite
    forwards but has an infinite derivative for signal_scale < 1.
  • Pin both with tests on mapper_util covering the NumPy value, the clean-warning case,
    NumPy/JAX parity with and without signal, and a finite jax.grad.
  • Sweep the package for the same where(cond, a / b, ...) pattern and record the result.
Detailed implementation plan

Affected Repositories

  • PyAutoArray (primary, and only)

The prompt lists autolens_workspace_test as a second repo, but its own body records the
script as already corrected to adapt to the lensed true image. Confirmed library-only.

Branch Survey

Repository Current Branch Dirty?
PyAutoArray main clean

No PyAutoArray claim in active.md / planned.md / parked.md; no matching remote
branch; no worktree conflict.

Branch: claude/autoarray-mapper-zero-signal-nan-jcck8q (a web-github session — no
gh, no task worktree, so the harness-designated branch rather than feature/...).

Implementation Steps

  1. autoarray/inversion/mappers/mapper_util.py, adaptive_pixel_signals_from step 7 —
    replace the guarded quotient with a safe denominator:

    max_sig = xp.max(pixel_signals)
    max_sig = xp.where(max_sig > 0, max_sig, 1.0)
    pixel_signals = pixel_signals / max_sig
  2. Same function, step 8 — exponentiate a safe base and select zero for non-positive
    signals afterwards, so the line stays finite and differentiable:

    zero_signal = pixel_signals <= 0.0
    safe_signals = xp.where(zero_signal, 1.0, pixel_signals)
    return xp.where(zero_signal, 0.0, safe_signals ** signal_scale)
  3. test_autoarray/inversion/pixelization/mappers/test_mapper_util.py — eight tests over a
    shared _zero_signal_kwargs() helper: NumPy finiteness and value, a
    RuntimeWarning-as-error assertion, finiteness parametrised over
    signal_scale in {0.5, 1.0, 2.0}, NumPy/JAX parity both with and without signal, and a
    finite jax.grad. JAX is an [optional] extra, so the JAX legs carry the repo's
    existing requires_jax skipif convention (as in test_delaunay.py).

Behaviour

Unchanged wherever it was previously well-defined. Two degenerate corners change
deliberately, both away from a wrong answer:

  • signal_scale == 0 with a zero signal returned 1.0 (the 0 ** 0 convention) and now
    returns 0.0.
  • A negative pixel signal returned NaN for a fractional signal_scale, or a spurious
    positive weight for an even integer one, and now returns 0.0 — the signals are
    documented as varying between 0 and 1, so a non-positive signal contributes nothing.
    (The sibling image_mesh/abstract_weighted.py already takes np.abs(adapt_data), so
    treating negative adapt signal as carrying no positive weight is consistent with the
    surrounding code.)

Pattern sweep (ask 3)

Done as an AST pass over the whole package rather than grep, so multi-line calls could not
slip past: every where(cond, x, y) with a division anywhere inside either branch.

Site Verdict
inversion/mappers/mapper_util.py:84 the bug — fixed here
inversion/mesh/interpolator/delaunay.py:390 safe — divides by the literal 3.0
inversion/mesh/interpolator/sibson.py:826 safe — divides by the literal 3.0
inversion/mesh/image_mesh/abstract_weighted.py:72 safe — explicit if max_value <= 0.0 early return
dataset/preprocess.py x3, dataset/imaging/dataset.py:507 NumPy-only preprocessing; divide by a user-supplied scalar, not on the JAX path
fit/fit_util.py:251, 452, 474 same bug class, live — filed separately (see below)

fit/fit_util.py:251 (chi_squared_map_with_mask_from) is on the JAX likelihood-gradient
path and was confirmed by direct reproduction: with a masked-out pixel carrying zero noise,
the forward value is finite (2.0) while jax.grad returns [2., 1., nan]. It is a
different module with its own test surface, so it is filed as its own PyAutoMind prompt
rather than widening this PR.

Key Files

  • autoarray/inversion/mappers/mapper_util.py — the fix, both steps 7 and 8
  • test_autoarray/inversion/pixelization/mappers/test_mapper_util.py — the eight tests

Testing

test_autoarray/inversion/pixelization/mappers/ — 27 passed. The two new legs that the
pre-fix code cannot pass (the clean-warning assertion and the finite-gradient assertion)
were verified red before the fix. test_autoarray/inversion test_autoarray/fit — 8
failures, all of them present identically on clean main (test_abstract.py sparse-operator
and test_factory.py imaging tests); no regression from this change.

Original Prompt

Click to expand starting prompt

Adapt-density mapper: a zero-signal adapt image is NaN on the JAX path and finite on NumPy

Type: bug
Target: PyAutoArray
Repos:

  • PyAutoArray
  • autolens_workspace_test
    Difficulty: small
    Autonomy: safe
    Priority: medium
    Status: formalised
    Filed: 2026-09-06

Found during the phase-6 rebuild of the ci-timing-fast-tests epic
(autolens_workspace_test#293 / #294). In
scripts/interferometer/jax_likelihood/rectangular.py an adapt image was built
from the unlensed true source profile evaluated on the image-plane grid — a
compact blob sitting where the Einstein ring is not, so under
RectangularRTUAdaptDensity almost every source pixel receives zero adapt
signal. The NumPy path returns a finite likelihood with a warning; the JAX path
returns NaN, and fitness._vmap collapses to the resample_figure_of_merit:

NumPy fit.log_likelihood: -3154.8962799401297
  .../autoarray/inversion/mappers/mapper_util.py:84: RuntimeWarning: invalid value
  encountered in divide
    pixel_signals = xp.where(max_sig > 0, pixel_signals / max_sig, pixel_signals)
JAX(no jit) fit.log_likelihood: nan
raw vmap: [-1.e+99]

xp.where(max_sig > 0, pixel_signals / max_sig, pixel_signals) guards the
selection but still evaluates the division; NumPy's 0/0 warning is discarded
by the selection while on the JAX path the NaN propagates into the likelihood
(the usual where-inside-grad/NaN-propagation pattern; a safe denominator
xp.where(max_sig > 0, max_sig, 1.0) is the standard fix). The workspace script
was corrected to adapt to the lensed true image (the right object anyway), so
nothing is red — but the two backends disagree on the same input, which is the
class of divergence the _test repos exist to catch.

Ask: (1) reproduce with a unit test on mapper_util where max_sig == 0 for
some source pixels, both backends; (2) fix the division so both paths agree
(finite, and identical); (3) check the same where(cond, a / b, a) pattern
elsewhere in autoarray.inversion (grep for / max_ and xp.where( around
divisions).

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions