fix(train): resume at exact global step instead of flooring to grad-accum boundary - #43
Open
longzhenren wants to merge 1 commit into
Open
longzhenren wants to merge 1 commit into
longzhenren wants to merge 1 commit into
Conversation
…ccum boundary compute_resume_position floored skip (and pulled aligned_global_step back) to a grad-accum boundary whenever global_step % grad_accum != 0. The premise was that a resumed run must start a fresh accumulation cycle, but that premise is false: Accelerate's load_state restores the accumulation counter, so resume already continues from the exact checkpointed position. Empirically (CPU, accelerate 1.12): running 3 micro-batches under accelerator.accumulate then save_state/load_state round-trips accelerator.step (3 -> 3), and save_state writes self.step while load_state assigns it back. With grad_accum=4, batches_per_epoch=10, resume at global_step=12 the old code returned (epoch=1, skip=0, aligned=10): epoch-1 batches 0-1 were re-fed, their gradients counted into the next optimizer step, and the per-step seed (keyed on aligned global_step) was wrong. The same floor also corrupted the derived opt_step (10//4=2 vs the checkpoint's real 12//4=3), shifting the LR schedule. Fix: skip = global_step % batches_per_epoch, aligned = global_step (drop the floor). grad_accum stays in the signature but is a no-op; grad_accum=1 behavior is unchanged. Tests updated to assert the exact resume position, with a new mid-accumulation case (ga=4, bpe=10, G=12 -> skip=2, aligned=12). Validated: tests/test_checkpointing.py 23 passed; train-adjacent subset (checkpointing, optimizer groups, training utils, prepare/unwrap, seeding, dataloader seed) 59 passed; ruff clean.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
compute_resume_positionfloorsskipto the previous grad-accum boundary whenglobal_step % grad_accum != 0, and pullsaligned_global_stepback to thatboundary. With grad_accum=4, batches_per_epoch=10, resuming at global_step=12 it
returns
(epoch=1, skip=0, aligned=10)instead of(1, 2, 12).The premise for the floor is that a resumed run must start a fresh accumulation
cycle. That premise does not hold: Accelerate's
load_staterestores theaccumulation counter, so resume already continues from the exact checkpointed
position.
Root cause
Verified empirically (CPU, accelerate 1.12): after 3 micro-batches under
accelerator.accumulate,save_state/load_stateround-tripsaccelerator.step(3 -> 3).save_statewritesself.stepandload_stateassigns it back. The training loop resumes through
accelerator.load_stateinload_full_state.With the floor, epoch-1 batches 0-1 are re-fed after resume, their gradients are
counted into the next optimizer step, and the per-step seed (keyed on aligned
global_step) shifts. The same floor also corrupts the derived
opt_step(10//4=2 vs the checkpoint's real 12//4=3), shifting the LR schedule.
Fix
skip = global_step % batches_per_epoch,aligned_global_step = global_step(drop the floor).
grad_accumstays in the signature but is a no-op.case (ga=4, bpe=10, G=12 -> skip=2, aligned=12).
Validation
pytest tests/test_checkpointing.py: 23 passed.prepare/unwrap, seeding, dataloader seed): 59 passed.
ruff check: clean.