From e54123b49471ac987660050300a8ec82b6a13a3c Mon Sep 17 00:00:00 2001 From: Jammy2211 Date: Sun, 4 Oct 2026 21:08:26 +0100 Subject: [PATCH] perf: back off the fnnls warm-start memo on scattered evaluation streams The passive-set memo's fallback guard judges a seed only after the seeded solve has run, so on a scattered (iid) stream every other solve paid for a bad seed (autolens_profiling#332: alma Delaunay solve 2.17x slower memo-on). Add a per-key back-off in nnls_memo: after 2 consecutive seeded-solve fallbacks the next 1, 2, 4, ... (cap 32) would-be-seeded solves start dense (still refreshing the entry). Each backed-off dense solve also checks, for free, whether the seed it skipped would have passed the guard against it (symmetric difference with the final passive set == warm_start_errors); if so the remaining skips are cancelled, so a stream that turns local regains the memo after one solve. An accepted real seed resets the streak. A local walk never falls back twice in a row, so its path and every reconstruction are unchanged. New stats key warm_start_backoff; memo_clear() clears entries and back-off state together. No config keys or tolerances change. Local witness (n=576, solver-only, min of 3, 3 seeds, 64 solves): iid on/off 1.49x -> 1.17x; walk 0.18x -> 0.18x; iid->walk phase recovery 0.69x. Refs PyAutoLabs/PyAutoArray#613 Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01S11WE9oj7Mvkfhc4EPBnyN --- .../inversion/inversion/inversion_util.py | 46 +++- autoarray/inversion/inversion/nnls_memo.py | 135 ++++++++++ .../inversion/inversion/test_nnls_memo.py | 240 ++++++++++++++++++ test_autoarray/util/test_jax_nnls.py | 7 +- 4 files changed, 419 insertions(+), 9 deletions(-) diff --git a/autoarray/inversion/inversion/inversion_util.py b/autoarray/inversion/inversion/inversion_util.py index 76454a552..abdcd6766 100644 --- a/autoarray/inversion/inversion/inversion_util.py +++ b/autoarray/inversion/inversion/inversion_util.py @@ -385,8 +385,9 @@ def reconstruction_positive_only_from( produced the returned reconstruction started from a memo seed, ``"dense"`` if it started from the sign of the unconstrained dense solve) and ``warm_start_fallback`` (`True` if a memo seed breached `Settings.nnls_warm_start_error_tolerance` and its entry was dropped, so the next solve for that key - restarts dense). They are set after `fnnls_cholesky` returns, so a diagnostic wrapping the solver - must read the dict it handed in *after* the evaluation, not at the point the solver returns. + restarts dense), and ``warm_start_backoff`` (`True` if the memo seed was skipped because the key's + seeds kept falling back -- see the back-off in `nnls_memo`). They are set after `fnnls_cholesky` + returns, so a diagnostic wrapping the solver must read the dict it handed in *after* the evaluation, not at the point the solver returns. Returns ------- @@ -564,8 +565,19 @@ def _record_pdip(converged, iterations): entry = nnls_memo.passive_set_get(key=key, n=n) if use_memo else None + # A key whose seeds keep falling back is backed off (see `nnls_memo`): + # this solve starts dense, but still refreshes the entry below. Only a + # solve that would otherwise have been seeded consumes a skip. + backoff = entry is not None and nnls_memo.backoff_should_skip(key=key) + + skipped_entry = entry if backoff else None + + if backoff: + entry = None + stats["seed_source"] = "dense" stats["warm_start_fallback"] = False + stats["warm_start_backoff"] = backoff if entry is not None: try: @@ -600,6 +612,12 @@ def _record_pdip(converged, iterations): if use_memo: error_fraction = stats["warm_start_errors"] / max(n, 1) + tolerance = settings.nnls_warm_start_error_tolerance + + guard_active = ( + tolerance is not None and np.isfinite(tolerance) and tolerance > 0.0 + ) + if stats["seed_source"] == "memo": # A memo seed is judged against the dense-sign start it # replaced, not against an absolute error count: the absolute @@ -609,19 +627,15 @@ def _record_pdip(converged, iterations): # next solve for this key restarts dense and refreshes the # reference -- no stale seed can be dragged through a run in a # regime the reference was never measured in. - tolerance = settings.nnls_warm_start_error_tolerance - - guard_active = ( - tolerance is not None and np.isfinite(tolerance) and tolerance > 0.0 - ) - if ( guard_active and error_fraction > tolerance * entry.dense_error_fraction ): nnls_memo.memo_drop(key=key) + nnls_memo.backoff_record_fallback(key=key) stats["warm_start_fallback"] = True else: + nnls_memo.backoff_record_accept(key=key) # The reference describes the dense-sign start, so it is # carried forward unchanged; only a dense solve refreshes it. nnls_memo.passive_set_put( @@ -630,6 +644,22 @@ def _record_pdip(converged, iterations): dense_error_fraction=entry.dense_error_fraction, ) else: + if ( + skipped_entry is not None + and guard_active + and nnls_memo.seed_error_fraction( + seed_passive_set=skipped_entry.passive_set, + passive_set=stats["passive_set"], + n=n, + ) + <= tolerance * error_fraction + ): + # The seed this backed-off solve skipped would have passed + # the guard against it, so the stream has turned local + # again: cancel the remaining skips and seed the next + # solve (whose own verdict decides whether the streak ends). + nnls_memo.backoff_end_skip(key=key) + nnls_memo.passive_set_put( key=key, passive_set=stats["passive_set"], diff --git a/autoarray/inversion/inversion/nnls_memo.py b/autoarray/inversion/inversion/nnls_memo.py index ef4f919c1..81cc4528e 100644 --- a/autoarray/inversion/inversion/nnls_memo.py +++ b/autoarray/inversion/inversion/nnls_memo.py @@ -37,6 +37,30 @@ # (`memo_drop`) and the next solve for that key restarts dense, refreshing the # reference. # +# That guard judges a seed only AFTER the seeded solve has run. On a +# scattered evaluation stream (iid draws, e.g. a nested sampler's early live +# points) every seed is bad, so the stream alternates dense solve -> bad +# seeded solve -> drop -> dense solve ...: half the solves pay for a bad seed. +# autolens_profiling#332 measured the alma interferometer Delaunay solve 2.17x +# slower memo-on than memo-off on such a stream. The per-key back-off below +# stops that: after `_NNLS_BACKOFF_AFTER_FALLBACKS` seeded solves for a key +# have fallen back in a row, the next solve(s) for it skip the seed and start +# dense (still refreshing the entry), for 1, 2, 4, ... up to +# `_NNLS_BACKOFF_MAX_SKIP` solves, then the seed is probed again. One accepted +# seed resets the streak, and a stream that never falls back twice in a row (a +# local walk) never meets the back-off at all. +# +# A backed-off solve also judges, for free, the seed it skipped: the dense +# solve's final passive set is the unique optimum's, so the seed's error count +# on this system is just the size of its disagreement with that set +# (`seed_error_fraction`) -- exactly what `warm_start_errors` would have +# reported had the seed been used. If the skipped seed would have passed the +# fallback guard against this dense solve, the remaining skips are cancelled +# and the next solve is seeded again: a stream that turns local regains the +# memo after one solve, without waiting out the schedule. The streak is kept +# (only an accepted REAL seeded solve resets it), so a shadow check that passes +# by chance on a scattered stream costs one probe and lengthens the next skip. +# # Disable with AUTOARRAY_NNLS_WARM_START=0. @@ -55,6 +79,31 @@ class MemoEntry(NamedTuple): _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES = 8 +# Consecutive seeded-solve fallbacks for one key before the back-off engages. +# 2, not 1: a single fallback already costs only one bad solve (the next solve +# restarts dense anyway), and a local walk that crosses one sharp change must +# keep the memo on the very next probe. +_NNLS_BACKOFF_AFTER_FALLBACKS = 2 + +# Cap on the number of solves skipped between probes of a backed-off key. On a +# stream where every seed is bad, one probe in (cap + 1) solves still pays for +# a bad seed; on a stream that turns local the memo is back within cap solves. +_NNLS_BACKOFF_MAX_SKIP = 32 + + +class BackoffState(NamedTuple): + """ + Per-key back-off bookkeeping: how many seeded solves in a row fell back + (`fallback_streak`) and how many upcoming solves still skip the seed + (`skip_remaining`). + """ + + fallback_streak: int + skip_remaining: int + + +_nnls_backoff: Dict[str, BackoffState] = {} + def memo_enabled() -> bool: """ @@ -121,3 +170,89 @@ def memo_drop(key: str) -> None: and refreshes the reference error fraction. A no-op if the key is absent. """ _nnls_passive_set_memo.pop(key, None) + + +def memo_clear() -> None: + """ + Forget every memo entry AND every back-off state, returning the process to + a cold start. Harnesses that compare memo arms must use this rather than + clearing `_nnls_passive_set_memo` alone: back-off state left by one arm + (e.g. an iid stream) would otherwise carry into the next. + """ + _nnls_passive_set_memo.clear() + _nnls_backoff.clear() + + +def backoff_should_skip(key: str) -> bool: + """ + Whether the solve for `key` about to run should skip the memo seed and + start dense, consuming one skip if so. False for a key with no back-off. + """ + state = _nnls_backoff.get(key) + + if state is None or state.skip_remaining <= 0: + return False + + _nnls_backoff[key] = state._replace(skip_remaining=state.skip_remaining - 1) + + return True + + +def backoff_record_fallback(key: str) -> None: + """ + Record that a seeded solve for `key` breached the fallback guard. Once + `_NNLS_BACKOFF_AFTER_FALLBACKS` have happened in a row, schedule the next + 1, 2, 4, ... (capped at `_NNLS_BACKOFF_MAX_SKIP`) solves to skip the seed. + """ + state = _nnls_backoff.get(key, BackoffState(0, 0)) + + streak = state.fallback_streak + 1 + + excess = streak - _NNLS_BACKOFF_AFTER_FALLBACKS + + skip = 0 if excess < 0 else min(2**excess, _NNLS_BACKOFF_MAX_SKIP) + + if ( + key not in _nnls_backoff + and len(_nnls_backoff) >= _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + ): + _nnls_backoff.pop(next(iter(_nnls_backoff))) + + _nnls_backoff[key] = BackoffState(fallback_streak=streak, skip_remaining=skip) + + +def backoff_end_skip(key: str) -> None: + """ + Cancel the remaining skips for `key` -- the next solve is seeded again -- + while keeping its fallback streak, so a failed probe resumes the back-off + at the escalated length. A no-op if the key has no back-off. + """ + state = _nnls_backoff.get(key) + + if state is not None: + _nnls_backoff[key] = state._replace(skip_remaining=0) + + +def backoff_record_accept(key: str) -> None: + """ + Record that a seeded solve for `key` passed the fallback guard, resetting + its back-off entirely. A no-op if the key has no back-off. + """ + _nnls_backoff.pop(key, None) + + +def seed_error_fraction(seed_passive_set: np.ndarray, passive_set, n: int) -> float: + """ + The fraction of a size-`n` solve's entries that the warm-start passive set + `seed_passive_set` gets wrong relative to the solve's final `passive_set` + -- the quantity `fnnls_cholesky` reports as `warm_start_errors / n` when it + is seeded from `seed_passive_set`. Used to judge a seed the back-off skipped + against the dense solve that ran instead. + """ + seed_mask = np.zeros(n, dtype=bool) + seed_mask[np.asarray(seed_passive_set, dtype=int)] = True + + final_mask = np.zeros(n, dtype=bool) + final_mask[np.asarray(passive_set, dtype=int)] = True + + return np.count_nonzero(seed_mask != final_mask) / max(n, 1) diff --git a/test_autoarray/inversion/inversion/test_nnls_memo.py b/test_autoarray/inversion/inversion/test_nnls_memo.py index f5f8a35fd..104f7d4bc 100644 --- a/test_autoarray/inversion/inversion/test_nnls_memo.py +++ b/test_autoarray/inversion/inversion/test_nnls_memo.py @@ -16,20 +16,31 @@ from autoarray.inversion.inversion import nnls_memo from autoarray.inversion.inversion.nnls_memo import ( + _NNLS_BACKOFF_AFTER_FALLBACKS, + _NNLS_BACKOFF_MAX_SKIP, _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES, + _nnls_backoff, _nnls_passive_set_memo, + backoff_end_skip, + backoff_record_accept, + backoff_record_fallback, + backoff_should_skip, memo_drop, + memo_clear, memo_key, passive_set_get, passive_set_put, + seed_error_fraction, ) @pytest.fixture(autouse=True) def _clean_memo(): _nnls_passive_set_memo.clear() + _nnls_backoff.clear() yield _nnls_passive_set_memo.clear() + _nnls_backoff.clear() def _normal_equations(seed, n=8, n_data=20): @@ -203,3 +214,232 @@ def test__memo_enabled__reads_the_environment(monkeypatch): def test__memo_key__separates_solve_sizes(): assert memo_key(n=3, fingerprint="mesh") != memo_key(n=4, fingerprint="mesh") + + + +# =================================================================== +# Per-key back-off on scattered evaluation streams (PyAutoArray#613) +# =================================================================== + + +def _skips_until_probe(key): + """How many consecutive would-be-seeded solves the back-off skips for `key`.""" + skipped = 0 + + while backoff_should_skip(key): + skipped += 1 + + return skipped + + +def test__backoff__engages_after_consecutive_fallbacks_and_grows_to_the_cap(): + assert _NNLS_BACKOFF_AFTER_FALLBACKS == 2 + + backoff_record_fallback("key") + assert _skips_until_probe("key") == 0 + + skips = [] + for _ in range(9): + backoff_record_fallback("key") + skips.append(_skips_until_probe("key")) + + assert skips == [1, 2, 4, 8, 16, 32, 32, 32, 32] + assert max(skips) == _NNLS_BACKOFF_MAX_SKIP + + +def test__backoff__accept_resets_the_streak_and_end_skip_keeps_it(): + for _ in range(3): + backoff_record_fallback("key") + + backoff_end_skip("key") + assert _skips_until_probe("key") == 0 + + # end_skip kept the streak: the next fallback escalates rather than restarting. + backoff_record_fallback("key") + assert _skips_until_probe("key") == 4 + + backoff_record_accept("key") + assert "key" not in _nnls_backoff + + backoff_record_fallback("key") + assert _skips_until_probe("key") == 0 + + # Unknown keys are no-ops. + backoff_end_skip("other") + backoff_record_accept("other") + assert backoff_should_skip("other") is False + + +def test__backoff__table_is_bounded(): + for i in range(_NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + 3): + backoff_record_fallback(f"key{i}") + + assert len(_nnls_backoff) == _NNLS_PASSIVE_SET_MEMO_MAX_ENTRIES + + +def test__memo_clear__forgets_entries_and_backoff(): + passive_set_put(key="key", passive_set=np.array([0]), dense_error_fraction=0.1) + for _ in range(3): + backoff_record_fallback("key") + + memo_clear() + + assert _nnls_passive_set_memo == {} + assert _nnls_backoff == {} + + +def test__seed_error_fraction__counts_the_symmetric_difference(): + assert seed_error_fraction(np.array([0, 2]), np.array([2, 0]), n=5) == 0.0 + assert seed_error_fraction(np.array([0, 1]), np.array([0, 2, 3]), n=5) == 0.6 + assert seed_error_fraction(np.array([], dtype=int), np.array([4]), n=5) == 0.2 + + +_N_STREAM = 40 + + +def _stream_system(Z, coeffs): + """Normal equations whose NNLS passive set follows the signs of `coeffs`.""" + x = Z @ coeffs + return Z.T @ Z + 0.1 * np.eye(Z.shape[1]), Z.T @ x + + +def _scattered_stream(seed, length=40): + """iid draws: every solve is an unrelated system, so no seed is any good.""" + rng = np.random.default_rng(seed) + return [ + _stream_system(rng.normal(size=(90, _N_STREAM)), rng.normal(size=_N_STREAM)) + for _ in range(length) + ] + + +def _walk_stream(seed, length=40, step=1e-3): + """A local walk: each system is a small perturbation of the previous one.""" + rng = np.random.default_rng(seed) + Z = rng.normal(size=(90, _N_STREAM)) + coeffs = rng.normal(size=_N_STREAM) + systems = [] + for _ in range(length): + systems.append(_stream_system(Z, coeffs)) + Z = Z + step * rng.normal(size=Z.shape) + coeffs = coeffs + step * rng.normal(size=coeffs.shape) + return systems + + +def _replay(monkeypatch, systems, memo=True): + """ + Solve `systems` in order through `reconstruction_positive_only_from` with one + memo key, returning the reconstructions and each solve's stats dict. + """ + import autoarray.util.fnnls as fnnls_mod + + original = fnnls_mod.fnnls_cholesky + captured = [] + + def _wrapped(ZTZ, ZTx, P_initial=np.zeros(0, dtype=int), stats=None, factor=None): + captured.append(stats) + return original(ZTZ, ZTx, P_initial, stats=stats, factor=factor) + + monkeypatch.setattr(fnnls_mod, "fnnls_cholesky", _wrapped) + + settings = aa.Settings( + use_positive_only_solver=True, + nnls_warm_start_memo=memo, + nnls_warm_start_error_tolerance=1.5, + ) + + reconstructions = [ + aa.util.inversion.reconstruction_positive_only_from( + data_vector=q, + curvature_reg_matrix=Q, + settings=settings, + fingerprint="stream", + ) + for Q, q in systems + ] + + monkeypatch.setattr(fnnls_mod, "fnnls_cholesky", original) + + # One stats dict per solve: the seeded attempt and the solve that returned + # share the dict (a raising seeded attempt never reaches the end). + stats = list({id(d): d for d in captured}.values()) + + return reconstructions, stats + + +@pytest.mark.parametrize("seed", [0, 1, 2, 3]) +def test__backoff__scattered_stream_stops_reseeding_and_keeps_the_answer( + monkeypatch, seed +): + systems = _scattered_stream(seed) + + expected, _ = _replay(monkeypatch, systems, memo=False) + + def bad_seeds(stats): + return sum(s["warm_start_fallback"] for s in stats) + + # Reference: the same replay with the back-off switched off. Every other + # solve is seeded and (bar a chance hit) every seed falls back. + with monkeypatch.context() as m: + m.setattr(nnls_memo, "_NNLS_BACKOFF_AFTER_FALLBACKS", 10**9) + _, stats_without = _replay(monkeypatch, systems) + + _nnls_passive_set_memo.clear() + _nnls_backoff.clear() + + reconstructions, stats = _replay(monkeypatch, systems) + + for reconstruction, reference in zip(reconstructions, expected): + assert reconstruction == pytest.approx(reference, rel=1e-9, abs=1e-11) + + assert bad_seeds(stats_without) >= 18 + assert not any(s["warm_start_backoff"] for s in stats_without) + + # With the back-off the probes thin out exponentially (solves 1, 3, 6, 10, + # 16, 26 on a stream with no chance hits): far fewer bad seeds are paid for. + assert bad_seeds(stats) <= 0.6 * bad_seeds(stats_without) + assert sum(s["warm_start_backoff"] for s in stats) >= 10 + + +@pytest.mark.parametrize("seed", [0, 1, 2, 3]) +def test__backoff__local_walk_is_untouched_bit_for_bit(monkeypatch, seed): + systems = _walk_stream(seed) + + reconstructions, stats = _replay(monkeypatch, systems) + + # The walk is what the memo is for: every solve after the first is seeded, + # no seed falls back, and the back-off never engages. + assert [s["seed_source"] for s in stats] == ["dense"] + ["memo"] * ( + len(systems) - 1 + ) + assert not any(s["warm_start_fallback"] for s in stats) + assert not any(s["warm_start_backoff"] for s in stats) + assert _nnls_backoff == {} + + # Bit-identical to the same replay with the back-off switched off entirely. + _nnls_passive_set_memo.clear() + monkeypatch.setattr(nnls_memo, "_NNLS_BACKOFF_AFTER_FALLBACKS", 10**9) + + reference, _ = _replay(monkeypatch, systems) + + for reconstruction, expected in zip(reconstructions, reference): + assert np.array_equal(reconstruction, expected) + + +@pytest.mark.parametrize("seed", [0, 1, 2]) +def test__backoff__scattered_then_local_regains_the_memo_after_one_solve( + monkeypatch, seed +): + scattered = _scattered_stream(seed, length=30) + systems = scattered + _walk_stream(seed + 100, length=20) + + # The walk starts from an unrelated system, so its first solve is dense + # (backed off or fallen back); the free check on that dense solve sees the + # skipped seed would have passed, and from the second walk solve on every + # solve is seeded again -- the long skip scheduled by the scattered phase is + # not waited out. + reconstructions, stats = _replay(monkeypatch, systems) + + walk_stats = stats[len(scattered) :] + + assert all(s["seed_source"] == "memo" for s in walk_stats[2:]) + assert not any(s["warm_start_fallback"] for s in walk_stats[2:]) diff --git a/test_autoarray/util/test_jax_nnls.py b/test_autoarray/util/test_jax_nnls.py index 9dac66580..a85024c8c 100644 --- a/test_autoarray/util/test_jax_nnls.py +++ b/test_autoarray/util/test_jax_nnls.py @@ -60,11 +60,16 @@ def test__reconstruction_positive_only_from__numpy_path_ignores_knobs(): @pytest.fixture(autouse=True) def _clear_nnls_memo(): - from autoarray.inversion.inversion.nnls_memo import _nnls_passive_set_memo + from autoarray.inversion.inversion.nnls_memo import ( + _nnls_backoff, + _nnls_passive_set_memo, + ) _nnls_passive_set_memo.clear() + _nnls_backoff.clear() yield _nnls_passive_set_memo.clear() + _nnls_backoff.clear() def _small_positive_only_system():