Skip to content

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

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

longzhenren wants to merge 4 commits into
OpenWAM-Official:mainfrom
longzhenren:fix/resume-exact-step

Conversation

@longzhenren

@longzhenren longzhenren commented Sep 30, 2026 •

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.

Resume semantics (boundary-only)

Sync boundaries are the 1-indexed epoch position hitting the accumulation quota or the epoch end (the prepared dataloader runs with Accelerate's default sync_with_dataloader=True, so the epoch end flushes the accumulation counter; measured boundaries for bpe=10/ga=4: {4, 8, 10, 14, 18, 20, ...}). The writer defers the resumable full state to those boundaries, discovery walks finished states newest-first past any legacy mid-cycle snapshot, and a mid-cycle checkpoint raises on resume (its pending micro-batch gradients were never persisted and cannot be reconstructed after load_state).

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

@wayrise wayrise left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Reviewed commit 6a28e7c60805377e94b5f90c5f6f2448dd8cba5b in an isolated checkout. Requesting changes because the claimed exact mid-accumulation resume still loses pending gradient contributions; restoring Accelerator.step alone is insufficient. See the inline finding.

Validation:

  • 102 training-related tests passed, including checkpointing, trainer, optimizer groups, scheduling, seeding and prepare/unwrap coverage.
  • Repository Ruff checks, training-source compilation and diff whitespace checks passed.
  • GPU comparisons used PyTorch 2.7.1+cu128, Accelerate 1.14.0 and DeepSpeed 0.18.9 on H200.
  • A tiny-model probe running the actual OpenWAMTrainer.train() loop diverged after checkpoints at microsteps 3 and 12 (grad_accum=4, 10 batches/epoch), while checkpoints at update boundaries 4 and 14 preserved the parameter trajectory.
  • A separate fresh-process probe using the production _build_accelerator configuration (bf16, ZeRO-2) reproduced the failure; its boundary-checkpoint control matched all 12 subsequent steps exactly.

The previous flooring logic is also not a complete solution, so this needs a consistent checkpoint/resume strategy rather than simply reverting the arithmetic. Please add a real save/load continuation regression test that compares parameters and optimizer state, not only returned positions or the accumulation counter.

Comment thread openwam/train/utils/checkpointing.py Outdated
The full-state snapshot persists the model, optimizer, and the Accelerate
step counter, but not the pending micro-batch gradients of an in-flight
accumulation cycle: saves run after every backward without requiring
accelerator.sync_gradients. Resuming a mid-cycle checkpoint exactly would
skip the micro-batches whose gradients were already accumulated but never
applied, so the next optimizer step sees only the tail micro-batch.

Mid-cycle checkpoints (global_step % grad_accum != 0) now degrade to the
last global sync boundary with a warning; boundary-aligned checkpoints
resume exactly. The alignment check uses the global step, so an in-epoch
skip that is not a grad_accum multiple no longer falsely degrades a
boundary resume. Docstring and tests now assert the guarantee the
save/load path actually provides.
@longzhenren

Copy link
Copy Markdown
Author

Confirmed: the pending-gradient loss is real, and your DeepSpeed bf16/ZeRO-2 numbers match what the save path allows (saves run after every backward whenever global_step % save_steps == 0, with no sync_gradients requirement, so a checkpoint after microbatch 3 of 4 carries no recoverable accumulation state).

Went with restricting exact resume to real sync boundaries (option b): persisting pending gradients under ZeRO sharding is not tractable in this code path. compute_resume_position now detects a mid-cycle checkpoint (global_step % grad_accum != 0), degrades to the last global sync boundary, and logs a warning with both step numbers; boundary-aligned checkpoints resume exactly as before. The boundary check works on the global step rather than the in-epoch skip, so a boundary landing inside an epoch with bpe % grad_accum != 0 (your gs=12, bpe=10 case) is no longer falsely degraded. Docstring and tests now assert the guarantee the save/load path actually provides: mid-cycle degrades with a warning (including the epoch-boundary aliasing case), boundaries resume exactly.

Training-adjacent subset (checkpointing, trainer, optimizer groups, dataloader seed): 59 passed. Ruff clean.

@wayrise wayrise left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Re-reviewed 05206b9ea990fb41ae86086c0b2a22866200dfa5 in a new isolated checkout. The warning acknowledges the missing gradients, but the new fallback does not yet make resume safe. I am keeping this review at Request changes for the two inline findings:

  1. Rewinding the returned data position leaves the restored accumulation counter at the original checkpoint phase, causing premature optimizer updates.
  2. Global-step divisibility is not the real synchronization boundary: the current Accelerator configuration already flushes at the end of a prepared dataloader. The new logic therefore also regresses valid boundary checkpoints.

Validation on H200 with PyTorch 2.7.1+cu128, Accelerate 1.14.0 and DeepSpeed 0.18.9:

  • 104 training-related tests passed; repository Ruff, training-source compilation and diff whitespace checks passed.
  • Fresh-process bf16/ZeRO-2 probes reproduced incorrect update timing after mid-cycle checkpoints; the step-4 boundary control matched every remaining loss and parameter value.
  • Production Accelerator factory probes with both ZeRO-1 and ZeRO-2, 10 batches/epoch and accumulation=4, observed sync steps 4, 8, 10, 14, 18, 20.
  • The actual OpenWAMTrainer.train() loop diverged in parameters after resuming from steps 3, 10, 12 and 14. The step-10 and step-14 cases also added a scheduler update; step 4 remained a passing control.

A concrete way to implement the proposed boundary-only strategy is to defer full-state saves until an actual synchronization boundary, record trustworthy boundary metadata, and reject unsupported existing mid-cycle checkpoints. If replay is supported instead, it needs a consistent restoration of the accumulation phase and training state. Please add real save/load continuation tests covering mid-cycle checkpoints and non-divisible epoch lengths; integer-only tests currently encode incorrect boundary assumptions.

Comment thread openwam/train/openwam_trainer.py
Comment thread openwam/train/utils/checkpointing.py Outdated
Degrading a mid-cycle checkpoint to an earlier boundary does not restore a
consistent accumulation phase: load_full_state has already set the
Accelerator step counter to the checkpointed value, so the first replayed
batch completes the loaded cycle and fires an optimizer step with one
batch worth of gradients treated as a full one (updates at reported steps
1, 5, 9 instead of 4, 8, 12).

The sync boundary itself was also wrong: the prepared dataloader runs with
Accelerate's default sync_with_dataloader=True, so every epoch end flushes
and resets the accumulation counter. With bpe=10, ga=4, updates fire at
global steps {4, 8, 10, 14, 18, 20, ...} — gs=10 is a real boundary
despite 10 % 4 != 0, and gs=12 is mid-cycle despite 12 % 4 == 0.

compute_resume_position now derives the boundary from the 1-indexed
position within the epoch (quota hit or epoch end) and raises ValueError
on mid-cycle checkpoints: pending micro-batch gradients are not persisted
and cannot be reconstructed after load_state, so the only safe resume is
from the nearest boundary checkpoint. Docstring and tests updated,
including the measured boundary set.
@longzhenren

Copy link
Copy Markdown
Author

Both P1s confirmed, redone along the lines they force.

P1-A (degradation does not restore the phase): correct, and the load_state counter makes any rewind unfixable — after load_full_state sets the Accelerator step to the checkpointed value, the first replayed batch completes the loaded cycle and fires an optimizer step with one batch of gradients treated as four (updates at reported steps 1, 5, 9). Dropping the degrade path entirely: compute_resume_position now raises ValueError for mid-cycle checkpoints (the full state does not persist pending micro-batch gradients and they cannot be reconstructed after load_state), so the only safe resume is from the nearest boundary checkpoint.

P1-B (epoch-end flush): confirmed — the prepared dataloader runs with Accelerate's default sync_with_dataloader=True, so the epoch end flushes and resets the counter. compute_resume_position now derives the boundary from the 1-indexed position within the epoch: sync boundary iff local % grad_accum == 0 or local == batches_per_epoch. With bpe=10, ga=4 that yields exactly the measured set {4, 8, 10, 14, 18, 20, ...}: gs=10 and gs=14 resume exact, gs=12 and gs=16 reject. The no-epoch-flush assumption is gone from the docstring, and the tests now assert the measured boundary set, the reject cases, and that grad_accum=1 never rejects.

Training-adjacent subset (checkpointing, trainer, optimizer groups, dataloader seed): 61 passed. Ruff clean.

@wayrise wayrise left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Re-reviewed 099c0f08f18b0461c8ea90dabf7be00782c714dc. The previous P1 correctness findings are addressed: real epoch-end boundaries are accepted, and unsupported mid-cycle checkpoints are rejected before further training.

Verification:

  • 106 training-related tests passed; repository Ruff, training-source compilation and diff whitespace checks passed.
  • Fresh-process H200 / bf16 / DeepSpeed ZeRO-2 comparisons at checkpoints 4, 10 and 14 (10 batches/epoch, accumulation=4) matched uninterrupted training at every subsequent microbatch: loss, BF16 weights, FP32 master weights, complete AdamW state, cosine-scheduler state and effective synchronization state.
  • Checkpoints 3 and 12 correctly raised ValueError before processing another batch.
  • The actual trainer loop independently confirmed the safe-boundary continuations and mid-cycle rejections.

I am retaining Request changes for one P2 integration issue, rather than the previous numerical-correctness issues: the save/discovery/retention workflow can still delete the last accepted checkpoint and automatically select a newly written checkpoint that this change rejects. This was reproduced through the actual trainer save/prune flow with both the default retention of 1 and retention of 3; details are inline.

Please keep the rejection of unsafe snapshots, but make future resumable saves and checkpoint selection preserve a usable recovery path, with an integration regression test spanning save, retention and resume. The PR description should also be updated to describe boundary-only resume and mid-cycle rejection.

Comment thread openwam/train/utils/checkpointing.py Outdated
Comment on lines +355 to +356
if grad_accum > 1 and local_step % grad_accum != 0 and local_step != batches_per_epoch:
raise ValueError(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[P2] Preserve a usable checkpoint before rejecting the automatically selected latest state

Rejecting a checkpoint with missing gradients is correct, but the writer still saves full states at every save_steps interval and retention treats those states as resumable. With 10 batches/epoch, accumulation=4, save_steps=4, save_full_states_for_resume=true and the default keep_last_k_ckpts=1, the real trainer saves a safe checkpoint at step 8, then saves an unsafe checkpoint at step 12 and deletes step 8. After an interruption, the only remaining state is step 12 and this new guard rejects it. The advice to use the nearest boundary checkpoint cannot be followed because it was pruned.

I also tested retention=3: states 4, 8 and 12 remain, but setup_output_dir/find_latest_accel_state still select 12 and fail; the normal resume_ckpt_path interface selects a run directory, not a specific older state. Thus simply increasing retention does not implement the advertised fallback.

Please coordinate this guard with saving/discovery/retention: for example, defer resumable full-state saves until an actual synchronization boundary and retain a valid recovery point, with a supported way to select a retained safe state when legacy directories contain unsafe snapshots. Keep the fail-fast behavior for unsupported states. Add a save→prune→resume integration test; the new integer tests cannot catch this failure.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Rechecked on 099af408e9a51ceb7c943b94f18c0ba2d0ef5816: the fresh-run keep-last-1 case and the initial legacy fallback are fixed. A remaining P2 affects a mixed legacy directory when retention is reduced.

Reproduction through the actual trainer:

  • Start with legacy states 4, 8 (safe) and 12 (unsafe), retained by an older run with keep_last_k_ckpts=3.
  • Resume with keep_last_k_ckpts=1, 10 batches/epoch, accumulation=4 and save_steps=4.
  • The new discovery correctly falls back to and loads step 8.
  • At step 12, no new full state is written, but the old unsafe accel_state_step_12 still exists. manage_checkpoints(..., 1) deletes states 4 and 8 and retains only that unsafe state.
  • Interrupt after pruning: the second resume fails with No resumable sync-boundary state.

This is conditional on a mixed old directory and the stricter retention limit; I am not claiming the ordinary fresh-run keep-last-1 path still fails. Please make retention account for the candidates already identified as unsafe, so one of them cannot displace the only usable recovery state. Add a fallback → continued save/prune → second-resume test; testing initial candidate choice alone misses this.

…nsafe states on resume

The resume guard rejects mid-cycle snapshots, but the writer still produced
one every save_steps and retention pruned the safe ones — with keep_last_k=1
a mid-cycle save left no resumable state at all, and discovery still picked
the unsafe newest snapshot and failed.

Three coordinated changes: the writer defers the resumable full state to
real sync boundaries (is_sync_boundary, mid-cycle steps keep their weights
line only); discovery returns every finished candidate newest-first instead
of just the latest; resume_if_configured walks that list and takes the
newest state compute_resume_position accepts, covering legacy dirs that
already hold mid-cycle snapshots. Retention needs no change: mid-cycle
steps no longer produce accel_state dirs, so pruning keeps the newest
boundary state.
@longzhenren

Copy link
Copy Markdown
Author

Confirmed the integration contradiction, fixed with the three coordinated changes you outlined.

Save side: the writer now defers the resumable full state to real sync boundaries (is_sync_boundary — the same 1-indexed epoch-position formula the resume guard uses). Mid-cycle steps keep their weights line only, so unsafe snapshots never enter retention. With keep_last_k=1 the sequence save@8 (boundary) -> save@12 (mid-cycle) now prunes the step-8 weights but keeps the step-8 state, which stays resumable.

Discovery side: find_accel_state_candidates returns every finished state newest-first (same completion marker rule), and resume_if_configured walks that list taking the newest state compute_resume_position accepts — covering legacy run dirs that already hold mid-cycle snapshots. A rejected candidate costs a metadata read; if nothing passes, the error lists every rejection.

Retention: no change needed — mid-cycle steps no longer produce accel_state_step_* dirs, so pruning keeps the newest boundary state.

Integration regression tests added (save -> prune -> discovery -> resume): keep_last_k=1 recovery (state@8 survives, weights@12 kept, resume lands on 8) and the legacy dir case (states @12 mid-cycle + @8 boundary, discovery walks 12 -> rejected -> 8 chosen). Checkpointing suite 30 passed; training-adjacent subset 64 passed. Ruff clean. PR description updated for the boundary-only resume and mid-cycle rejection semantics.

@wayrise wayrise left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Re-reviewed 099af408e9a51ceb7c943b94f18c0ba2d0ef5816. The ordinary fresh-run retention case now passes: saving weights at step 12 leaves the safe step-8 full state available. The new legacy-candidate selection also works: unsafe step 12 is rejected before loading and only safe step 8 is loaded.

Request changes remains for two P2 save/retention issues:

  1. A mid-cycle save request is discarded, rather than serviced at the next synchronization boundary (inline finding below).
  2. After falling back in a legacy directory, reducing retention to 1 can still prune the chosen safe state in favor of the unsafe legacy state. I have continued the existing retention discussion with the precise precondition and a second-resume reproduction.

Validation:

  • 109 training-related tests passed; repository Ruff, compilation of training sources and diff whitespace checks passed.
  • H200 / bf16 / DeepSpeed ZeRO-2 fresh-process fallback from 12 to 8 matched the uninterrupted run for all 12 remaining microbatches, including loss, model/master weights, AdamW and cosine-scheduler state. An unsafe-only directory failed before loading any state.
  • The actual trainer loop reproduced both remaining issues, including a 5-epoch configuration interrupted at step 4500 with no resumable state.

Please exercise the real writer in the regression tests. The new tests manually create checkpoint directories and simulate candidate selection, so they do not verify deferred saving or fallback → continued training → pruning → a second resume.

Comment on lines 368 to +372
if save_steps and global_step > 0 and global_step % save_steps == 0:
save_weights(self.accelerator, self.architecture, output_path, global_step, final=False)
if save_full_states_for_resume:
save_full_state(self.accelerator, output_path, global_step, opt_step, epoch)
if is_sync_boundary(global_step, len(dataloader), grad_accum):
save_full_state(self.accelerator, output_path, global_step, opt_step, epoch)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[P2] Carry the save request to the next synchronization boundary

The full-state write is still nested under global_step % save_steps == 0. When a scheduled save is mid-cycle, this branch logs that it is deferred but records no pending request; the next sync boundary never enters the branch. Only an intersection of the save interval and synchronization schedule can produce a resumable checkpoint.

I reproduced this with the actual trainer loop, save_full_states_for_resume=true, 1001 batches/epoch, accumulation=4, the default save_steps=2000, and 5 epochs. Steps 2000 and 4000 were skipped; real sync boundaries 2001 and 4003 did not save either. Interrupting at step 4500 left only checkpoint_step_4000.safetensors and no accel_state_step_* directory. The run therefore cannot resume despite two scheduled saves.

Please retain a pending full-state-save request and service it at the next actual synchronization boundary, even when that step is not divisible by save_steps. A smaller regression case is 100 batches/epoch, accumulation=4, save interval=3: the request at 3 should create a resumable state at 4.

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.

2 participants