Skip to content

fix(train): resume at exact global step instead of flooring to grad-accum boundary - #43

Open
longzhenren wants to merge 1 commit into
OpenWAM-Official:mainfrom
longzhenren:fix/resume-exact-step
Open

longzhenren wants to merge 1 commit into
OpenWAM-Official:mainfrom
longzhenren:fix/resume-exact-step

Conversation

@longzhenren

Copy link
Copy Markdown

Problem

compute_resume_position floors skip to the previous grad-accum boundary when
global_step % grad_accum != 0, and pulls aligned_global_step back to that
boundary. 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_state restores the
accumulation 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_state round-trips
accelerator.step (3 -> 3). save_state writes self.step and load_state
assigns it back. The training loop resumes through
accelerator.load_state in load_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_accum stays in the signature but is a no-op.
  • grad_accum=1 behavior is unchanged.
  • Tests updated to assert the exact resume position, plus a new mid-accumulation
    case (ga=4, bpe=10, G=12 -> skip=2, aligned=12).

Validation

  • pytest tests/test_checkpointing.py: 23 passed.
  • Train-adjacent subset (checkpointing, optimizer groups, training utils,
    prepare/unwrap, seeding, dataloader seed): 59 passed.
  • ruff check: clean.

…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.
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.

1 participant