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
2 changes: 1 addition & 1 deletion bergson/magic/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -633,7 +633,7 @@ def train(

# Fast-forward the checkpoint schedule to where we're resuming from.
while next_save < start:
next_save = next_save_index(next_save, n, save_mode)
next_save = next_save_index(next_save, n, save_mode, save_interval)

pending_save: SaveFuture | None = None

Expand Down
53 changes: 53 additions & 0 deletions tests/test_magic.py
Original file line number Diff line number Diff line change
Expand Up @@ -868,6 +868,59 @@ def boom(i, loss):
torch.testing.assert_close(resumed_scores, fresh_scores, atol=1e-12, rtol=1e-6)


def test_magic_resume_interval_fast_forwards_schedule():
"""Resuming a ``save_mode="interval"`` run fast-forwards the checkpoint
schedule with ``save_interval``; the resume branch used to omit it and
raise ``save_mode='interval' requires save_interval > 0``."""
n = 9
crash_at = 5

trainer, fwd_state, model = _fresh_trainer()
stream = _multi_step_stream(n)
with tempfile.TemporaryDirectory() as ckpt_dir:
trainer.train(
fwd_state,
stream,
inplace=True,
save_dir=ckpt_dir,
save_mode="interval",
save_interval=3,
)
fresh_steps = _saved_steps(ckpt_dir)

with tempfile.TemporaryDirectory() as ckpt_dir:

def boom(i, loss):
if i == crash_at:
raise RuntimeError("simulated crash")

trainer, fwd_state, model = _fresh_trainer()
stream = _multi_step_stream(n)
with pytest.raises(RuntimeError, match="simulated crash"):
trainer.train(
fwd_state,
stream,
inplace=True,
save_dir=ckpt_dir,
save_mode="interval",
save_interval=3,
log_fn=boom,
)

trainer, fwd_state, model = _fresh_trainer()
stream = _multi_step_stream(n)
trainer.train(
fwd_state,
stream,
inplace=True,
save_dir=ckpt_dir,
save_mode="interval",
save_interval=3,
resume=True,
)
assert _saved_steps(ckpt_dir) == fresh_steps


def test_magic_resume(dataset):
"""Resume from a checkpoint mid-training and verify identical final state."""
device = "cpu"
Expand Down
Loading