Conversation
|
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
left a comment
There was a problem hiding this comment.
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
|
Ok that seems reasonable, I’ll take a look tomorrow |
There was a problem hiding this comment.
Looks good, open an issue @rz4 for the silent bug so we don't forget, I think we can merge this for now.
Summary
Adds
state_dict()/load_state_dict()toBaseDriftDetectorso 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, andsrc/main.pystill 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_ATTRStuple; 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, ...]]whereNonemeans "has not opted in" and()means "genuinely stateless" — so a custom detector that never declared its state raisesNotImplementedErrorinstead 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)EnsembleDetectoroverrides both to delegate to its sub-detectorsBreaking 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
NotImplementedErrorif checkpointed before declaring_STATE_ATTRS.Performance (optional)
Security & Privacy
Dependencies
None
Testing Plan
poetry run pytest tests/test_drift_detection.py -k Checkpointing -vDocumentation
Checklist
ruff format --checkruff check .mypy srcpytest -qRisk & 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.