Skip to content

fix: Keep masked fit utility gradients finite #604

Description

@Jammy2211

Overview

Prevent invalid divisions in discarded branches of three fit utilities. Return zero for zero-data residual fractions consistently in masked and unmasked helpers.

Plan

  • Pin the forward behavior and failing gradients before changing the utilities.
  • Guard denominators before division at all three sites.
  • Make zero-data residual fractions consistently return zero.
  • Validate NumPy results and compiled JAX gradients, then ship linked PRs.
Detailed implementation plan

Suggested branch: feature/fit-util-masked-division

Classification: Both; PyAutoArray primary, autolens_workspace_test companion for required permanent JAX regression coverage under repository policy. Small, independent.

  1. In autoarray/fit/fit_util.py, change chi_squared_map_with_mask_from, residual_flux_fraction_map_from, residual_flux_fraction_map_with_mask_from. Compute inclusion conditions first, replace excluded denominators with 1.0, divide, then select zero for excluded pixels.
  2. For masked residual fractions use (mask == 0) & (data != 0), matching the unmasked helper's zero-data semantics. Document this intentional forward change. Do not hide zero noise on an included chi-squared pixel.
  3. Extend test_autoarray/fit/test_fit_util.py with NumPy forward cases: masked zero denominators, included zero data, ordinary nonzero values, all-masked and mixed masks. Confirm existing defined results unchanged.
  4. Add scripts/misc/jax_assertions/fit_util_masked_division.py in autolens_workspace_test. For each site require finite jax.grad results and zero derivative on excluded entries, both eager and jitted; differentiate with respect to numerator and denominator where applicable. Require agreement with analytic included gradients and NumPy forwards. Include nonzero residual over zero data so the guard genuinely matters. Show failures on origin/main before accepting the repair.
  5. Run focused fit tests and full test_autoarray, new JAX script and applicable existing likelihood parity smoke. Register the regression with the workspace smoke runner using its existing conventions. Library PR first, companion workspace PR at the library-first gate.

Worktree root: /home/jammy/Code/PyAutoLabs/.worktrees/autoarray-bundle-1, created once using Brain's worktree helper with PYAUTO_WT_ROOT inside this workspace. Shared repo worktrees: PyAutoArray, PyAutoGalaxy, autolens_workspace_test. Each member starts from origin/main on its own branch; execute and ship sequentially before switching the shared PyAutoArray worktree. One primary PyAutoArray issue and registry entry per member, linked companion PRs per repository as explicitly authorized. No merges.

Branch survey: PyAutoArray, PyAutoGalaxy and autolens_workspace_test canonical checkouts are clean on main. No target repo claims in active.md. Heart reports an unregistered sparse-operator-oversampling-cache/PyAutoArray worktree with 11 dirty files: preserve it and obtain overlap acknowledgement before setup. Recent PyAutoArray branches: main, feature/sparse-operator-oversampling-cache, chore/session-start-hook-regen, feature/delaunay-area-magnification-audit, claude/autonerves-floor-regime-stamp. Recent PyAutoGalaxy branches: main, chore/session-start-hook-regen. Recent autolens_workspace_test branches: main, feature/point-audits-wheel-provenance, feature/point-solver-image-accuracy, feature/point-solver-duplicate-policy, chore/session-start-hook-regen.

Execution: one native Sol delegate per member, sequential within shared repositories. Parent owns judgment and lifecycle. Pass the approved issue plan, branch, exact starting commit, worktree, permitted files, validation requirements; stop and return exact failure evidence rather than weakening tests. Full logs stay in ignored scratch. Applicable full library suites and workspace smoke checks, authoritative Heart verdict, then ship separately. Report counts without inventing results; CI must check Python 3.12 and 3.13. No tests have yet run for this bundle.

Original Prompt

Click to expand starting prompt

fit_util: masked divisions NaN the gradient on the JAX path

Type: bug
Target: PyAutoArray
Repos:

  • PyAutoArray
    Difficulty: small
    Autonomy: safe
    Priority: medium
    Status: formalised
    Consequence: glance
    Witness: A jax.grad test on each of the three fit_util.py xp.where division sites (chi_squared_map_with_mask_from, residual_flux_fraction_map_from, residual_flux_fraction_map_with_mask_from) asserts a finite gradient with a masked-out pixel carrying zero noise or zero data, a forward test covers the unmasked-zero-data case at the _with_mask site, and forward NumPy and JAX values are unchanged.
    Review-minutes: 3
    Filed: 2026-09-10

Found by the ask-(3) pattern sweep of
complete/2026/09/autoarray-mapper-zero-signal-nan.md (PyAutoArray#548, shipped in
#549), and
split out of it on the human's call rather than widening that PR into a second
module. Same bug class, different module, and on the face of it the more serious
of the two: this one sits on the likelihood-gradient path.

autoarray/fit/fit_util.py has three xp.where(cond, <expr with a division>, …)
sites where the where guards only the selection — the division is still
evaluated for every element, so a zero denominator puts a NaN in the discarded
branch. Forwards it is invisible; under jax.grad it propagates.

fit_util.py:251  chi_squared_map_with_mask_from
                 xp.where(mask == 0, xp.square(residual_map / noise_map), 0)
fit_util.py:452  residual_flux_fraction_map_from
                 xp.where(data != 0, residual_map / data, 0)
fit_util.py:474  residual_flux_fraction_map_with_mask_from
                 xp.where(mask == 0, residual_map / data, 0)

Reproduced for :251, the chi-squared map, with a masked-out pixel carrying zero
noise (the ordinary case — noise maps are commonly zeroed outside the mask):

mask      = np.array([0, 0, 1])     # 0 == included
noise_map = np.array([1.0, 2.0, 0.0])
# forward : 2.0                     <- finite, so nothing catches it
# grad    : [ 2.,  1., nan]

:474 has a second problem on top of the shared one: it guards on the mask, not
on data, so a zero data value anywhere inside the unmasked region is a NaN
forwards, on both backends.

Ask: (1) pin each of the three with a jax.grad test asserting a finite
gradient, plus a forward test for the :474 unmasked-zero case; (2) fix with a
safe denominator — denom = xp.where(cond, denom, 1.0) then divide — the idiom
mapper_util.adaptive_pixel_signals_from now uses after #548, and that
pixel_counts in that same function used already; (3) decide whether :474
should guard on data != 0 as well as the mask, or whether a zero data value in
the unmasked region should stay loud.

Note for whoever picks this up: #548's investigation showed the forward NumPy
and JAX results agreed at its site, and the same is true here — do not expect a
forward-value divergence to reproduce it. The gradient is the witness.

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