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
46 changes: 38 additions & 8 deletions autoarray/inversion/inversion/inversion_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
-------
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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(
Expand All @@ -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"],
Expand Down
135 changes: 135 additions & 0 deletions autoarray/inversion/inversion/nnls_memo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.


Expand All @@ -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:
"""
Expand Down Expand Up @@ -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)
Loading
Loading