Skip to content

Zilinghan/checkpointing - #135

Open
Zilinghan wants to merge 2 commits into
mainfrom
zilinghan/checkpointing
Open

Zilinghan wants to merge 2 commits into
mainfrom
zilinghan/checkpointing

Conversation

@Zilinghan

Copy link
Copy Markdown
Collaborator

Summary

Adds state_dict() / load_state_dict() to BaseDriftDetector so a detector's accumulated stream state can be saved and restored. This is the first piece of resumable runs: the detector is the one component whose verdict depends on everything it has already seen, so a run that restarts it from scratch mid-stream produces different science, not just repeated work.

Motivation & Context

#126 added BaseModelHarness.save_ckpt(), but checkpoints are currently write-only — there is no load path anywhere, and src/main.py still carries a # TODO: Save a model checkpoint. Full run-resume needs model weights, optimizer moments, monitor counters, harness stream position, CL updater state (EWC Fisher, KFAC factors) and detector state. This PR does detector state only, deliberately, to keep the design decision reviewable on its own.

Approach

Each detector declares the attributes that evolve with the stream in a _STATE_ATTRS tuple; the base class implements save/load once against that declaration, so adding a detector is a one-line change rather than another pair of methods. Three deliberate choices: hyperparameters are excluded from the snapshot (they are rebuilt from config, so a stale checkpoint can never silently override a config someone edited); the river detector object is stored directly rather than replaying values (replay only works for ADWIN, the only detector that retains its value history, and is O(n) on resume — the trade-off is that checkpoints are coupled to the river version); and state is deep-copied on both save and load, so a snapshot never aliases the live detector and one checkpoint loaded into two detectors does not leave them sharing a mutable object.
The declaration is Optional[tuple[str, ...]] where None means "has not opted in" and () means "genuinely stateless" — so a custom detector that never declared its state raises NotImplementedError instead of silently checkpointing nothing.

Screenshots / Logs (optional)

API / CLI Changes

  • BaseDriftDetector.state_dict() -> Dict[str, Any] (new)
  • BaseDriftDetector.load_state_dict(state: Dict[str, Any]) -> None (new)
  • BaseDriftDetector._STATE_ATTRS: Optional[tuple[str, ...]] = None (new class attribute; subclasses declare their own)
  • EnsembleDetector overrides both to delegate to its sub-detectors
  • No config keys, CLI flags or env vars added

Breaking Changes

None. Both methods are concrete on the base class, not abstract, so existing detector subclasses — including any written outside this repo — continue to work unchanged; they simply raise a clear NotImplementedError if checkpointed before declaring _STATE_ATTRS.

Performance (optional)

Case Before After Notes
foo() 123 ms 88 ms Median of 50 runs

Security & Privacy

  • No secrets committed
  • Input validation added where needed

Dependencies

None

Testing Plan

  • Unit tests
  • Integration tests
  • e2e / smoke test
  • Manual steps: poetry run pytest tests/test_drift_detection.py -k Checkpointing -v

Documentation

  • Docstrings updated
  • User docs / README updated
  • CHANGELOG entry

Checklist

  • Code formatted (Ruff) → ruff format --check
  • Lint passes (Ruff) → ruff check .
  • Types pass (mypy/pyright) → mypy src
  • Tests pass (pytest) → pytest -q
  • Backward compatibility considered
  • Adequate comments for tricky parts
  • CI green

Risk & Rollback Plan

Probably not needed in the beginning

Low. The new code is additive and nothing in the runtime path calls it yet, so rollback is a plain revert with no migration or config change.

Notes for Reviewers

Note that nothing calls this API yet. Wiring it in requires the full run-state checkpoint (monitor counters, optimizer, harness stream position), which is the follow-up PR.

@anagainaru

Copy link
Copy Markdown
Collaborator

I am a bit confused by why we need this.

You save the checkpoints and you can restart by pointing to the checkpoint in the toml file. The stream of data shouldn't be saved since we are dealing with streams and not files (I know we are faking a stream by reading a file but eventually we will move to full stream).

What is the use case?

@rz4

rz4 commented Oct 1, 2026

Copy link
Copy Markdown
Collaborator

I am a bit confused by why we need this.

You save the checkpoints and you can restart by pointing to the checkpoint in the toml file. The stream of data shouldn't be saved since we are dealing with streams and not files (I know we are faking a stream by reading a file but eventually we will move to full stream).

What is the use case?

We currently checkpoint the state of the model weights at the end of the CL loop only. Like how Zilinghan describes in motivations and context, a full resume would require additional states of separate components of the framework. This PR is just for checkpointing the state of the detectors. The detectors accumulate over the stream, so instead of a re-playing a stream from the start for drift detection, this would let you get back to a state of the detector mid-stream.

@rz4 rz4 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good. There was a specific silent bug that may need to be fixed in a subsequent PR if there is a usecase for changing the detector hyper-parameters at a resume. Since this checkpointing is not exercised yet during drift detection, I don't see any remaining issues with approving this PR.

On resume, load_state_dict replaces the wrapped river object wholesale, so river-level hyper parameters silently keep their checkpointed values and ignore any changed config. The wrapper-level thresholds do pick up the new config values. Here's the problem summarized as a test function:

    @pytest.mark.parametrize(
        ("make_detector", "param", "old", "new"),
        [
            pytest.param(ADWINDetector, "delta", 0.002, 0.05, id="adwin"),
            pytest.param(KSWINDetector, "alpha", 0.005, 0.001, id="kswin"),
            pytest.param(
                PageHinkleyDetector, "threshold", 50.0, 10.0, id="page_hinkley"
            ),
        ],
    )
    def test_river_params_come_from_checkpoint_not_config(
        self, make_detector, param, old, new
    ):
        """Restoring replaces the river object wholesale, so
        river-level hyperparameters silently keep their checkpointed values
        even when the detector was rebuilt with new ones, while wrapper-level
        params (the thresholds) take the new config."""
        state = make_detector(**{param: old}, minor_threshold=0.3).state_dict()

        resumed = make_detector(**{param: new}, minor_threshold=0.9)
        resumed.load_state_dict(state)

        assert getattr(resumed.detector, param) == old  # stale checkpoint value
        assert resumed.minor_threshold == 0.9  # new config value
        # wrapper now disagrees with its own river object
        assert getattr(resumed, param) == new # new config value

@anagainaru

Copy link
Copy Markdown
Collaborator

Ok that seems reasonable, I’ll take a look tomorrow

@anagainaru anagainaru left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks good, open an issue @rz4 for the silent bug so we don't forget, I think we can merge this for now.

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.

3 participants