feat: add max_steps to all trainers - #382
stephantul wants to merge 1 commit into
Conversation
fit now accepts max_steps, which stops training after that many steps, also in the middle of an epoch. When it is reached, the model is validated one last time, unless that step was already a validation step, and the best checkpoint is kept. max_steps takes precedence over min_epochs.
Codecov Report✅ All modified and coverable lines are covered by tests.
🚀 New features to boost your workflow:
|
|
| def test_run_training_loop_stops_at_max_steps() -> None: | ||
| """Training stops after max_steps, across epochs, and validates once at the end.""" | ||
| loss = _run_counting_loop(max_steps=13) | ||
| assert loss.train_calls == 13 | ||
| assert loss.val_calls == 2 * 2 |
There was a problem hiding this comment.
Best checkpoint remains untested The new test checks training and validation call counts, but it never checks the returned checkpoint. It would pass even if training validated at
max_steps and then returned the final weights instead of an earlier, better checkpoint. A test with controlled validation results would protect that behavior.
Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!
fit now accepts max_steps, which stops training after that many steps, also in the middle of an epoch. When it is reached, the model is validated one last time, unless that step was already a validation step, and the best checkpoint is kept. max_steps takes precedence over min_epochs. This is mostly prep work for a trainer which doesn't have the concept of epochs.