Skip to content

Refuse degenerate fits in the AMICA wrapper (#50) - #54

Closed
neuromechanist wants to merge 2 commits into
51-multi-model-ng-log-likelihood-is-002-lower-and-more-variable-than-fortranfrom
50-amica-wrapper-marks-is_fitted_-even-on-a-degenerate-nan_llsingular_ll-fit
Closed

neuromechanist wants to merge 2 commits into
51-multi-model-ng-log-likelihood-is-002-lower-and-more-variable-than-fortranfrom
50-amica-wrapper-marks-is_fitted_-even-on-a-degenerate-nan_llsingular_ll-fit

Conversation

@neuromechanist

Copy link
Copy Markdown
Member

Closes #50. Stacked on #53 (issue #51) — review the last commit; the base retargets to main automatically when #53 merges.

Problem

AMICA.fit() set is_fitted_ = True unconditionally, so transform()/get_mixing_matrix()/get_unmixing_matrix() ran on a degenerate fit (stop_reason in nan_ll/singular_ll) and returned NaN sources with no exception, while state_dict()/save() already refused such a model. The n_models>1 work made this more reachable.

Fix — consistent, fail-loud contract (mirrors state_dict's refusal)

  • fit() sets is_fitted_ only when the fit converged, and exposes converged_ (bool) and stop_reason_ (str) for inspection; a degenerate fit logs a wrapper-level warning.
  • New _check_usable() guard raises a clear degenerate RuntimeError (naming the stop_reason) from transform/get_mixing_matrix/get_unmixing_matrix/save. An unfitted model still raises a distinct "must be fitted" ValueError.
  • load() carries stop_reason_/converged_ through (a saved model is always converged, since state_dict refuses degenerate ones).

Contract chosen per maintainer decision: expose attributes + refuse output (not a hard error at fit()), so a diverged run stays inspectable via stop_reason_/ll_history_.

Tests

3 new wrapper tests: converged_/stop_reason_ exposed on a normal fit; unfitted vs degenerate error messages are distinct and diagnosable; transform/get_*/save all refuse a degenerate model. The degenerate marker is forced after a real fit (the backend does not diverge on the clean sample EEG), the same pattern test_state_dict_refuses_degenerate_model already uses — real data, no mocks. 10 wrapper tests pass; ruff clean; ty no new diagnostics.

AMICA.fit() marked is_fitted_ = True unconditionally, so transform()/
get_mixing_matrix()/get_unmixing_matrix() ran on a degenerate fit
(stop_reason nan_ll/singular_ll) and returned NaN sources with no error,
while state_dict()/save() already refused such a model.

Make the wrapper's contract consistent:
- fit() sets is_fitted_ only when the fit converged, and exposes converged_
  (bool) and stop_reason_ (str) for inspection; a degenerate fit logs a
  wrapper-level warning.
- A new _check_usable() guard raises a clear degenerate error (naming the
  stop_reason) from transform/get_mixing/get_unmixing/save, mirroring
  state_dict()'s refusal. An unfitted model still raises a distinct
  "must be fitted" error.
- load() carries stop_reason_/converged_ through (a saved model is always
  converged, since state_dict refuses degenerate ones).

Tested: 3 new wrapper tests (converged/stop_reason exposed on a normal fit;
unfitted vs degenerate error messages; transform/get_*/save all refuse a
degenerate model) following the established real-fit-then-force-stop_reason
pattern; 10 wrapper tests pass. ruff clean; ty no new diagnostics.
PR review (3 Sonnet reviewers):
- CRITICAL (regression I introduced): _check_usable keyed on model_ is None,
  but fit() assigned self.model_ before training, so a mid-fit exception left a
  half-built backend that the guard let through (and, on refit, stale
  is_fitted_=True). fit() now trains a LOCAL backend and only publishes it to
  self on success -- a first-fit crash keeps model_ None (clean "not fitted"),
  a failed refit keeps the last good model.
- Strengthened the degenerate test to a REAL divergence: a single NaN injected
  into the real EEG forces an actual nan_ll stop (error-path robustness test,
  not a parity oracle), so it exercises fit()'s real is_fitted_=False/warning
  bookkeeping instead of a forced marker, and now also covers fit_transform.
- Added load() converged_/stop_reason_ round-trip assertions.

10 wrapper tests pass; ruff clean; ty no new diagnostics.
@neuromechanist

Copy link
Copy Markdown
Member Author

Review response (3 Sonnet reviewers: code, silent-failure, tests)

code-reviewer: no ≥80-confidence issues (ran the wrapper suite 10/10, verified lrate_floor is correctly treated as converged, no caller relies on the old is_fitted_ semantics).

Fixed (commit 7bdf38a)

  • CRITICAL — regression this PR introduced (silent-failure). _check_usable keys off model_ is None, but fit() assigned self.model_ before training, so an exception thrown out of backend.fit() (OOM, LinAlgError, interrupt, …) left a half-built backend that the guard let through into an opaque TypeError — and on a refit left is_fitted_/converged_ falsely True from the prior fit. fit() now trains a local backend and publishes it to self only after fit() returns: a first-fit crash keeps model_ is None (clean "not fitted"), a failed refit keeps the last known-good model. The old code happened to be safe here via is_fitted_; my initial AMICA wrapper marks is_fitted_ even on a degenerate (nan_ll/singular_ll) fit #50 change regressed it — good catch.
  • Real degenerate test (test-analyzer gaps Include a PyTorch implementation #1 + Internal function handling #3). Replaced the forced-stop_reason_ marker with a genuine divergence: a single NaN injected into the real EEG forces an actual nan_ll stop (an error-path robustness test, not a parity/correctness oracle, so it doesn't violate the no-synthetic-data policy). This now exercises the real fit() bookkeeping (is_fitted_=False, converged_=False, the wrapper logger.warning) and adds a fit_transform refusal check.
  • load() wiring (test-analyzer gap Handle and verify tests #2). test_ng_save_load_roundtrip now asserts converged_/stop_reason_ round-trip.

Not changed (with rationale)

  • The mid-fit-exception path is verified by inspection, not a new test. The fix is a pure reorder (assign self.* only after backend.fit() returns), and a real post-construction backend.fit() exception isn't inducible on the clean sample EEG without mock.patch (which the reviewer used to demo, but the NO-MOCK policy forbids in the suite). The model_ is None → "not fitted" outcome is covered by test_unfitted_output_raises_not_fitted.
  • _DEGENERATE_STOP_REASONS reached from the wrapper (code-reviewer/​~40-confidence nit). Left as-is: acceptable intra-package coupling that avoids duplicating the string list; a classmethod accessor would be churn for three same-author call sites.
  • converged_ naming for max_iter/lrate_floor stops: the docstring already defines it as "ended on a usable stop rather than a degenerate one," so it's precise as documented.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant