From 6a28e7c60805377e94b5f90c5f6f2448dd8cba5b Mon Sep 17 00:00:00 2001 From: "LongZhenren(Zhibo Zhang)" Date: Wed, 30 Sep 2026 18:00:47 +0800 Subject: [PATCH 1/4] fix(train): resume at exact global step instead of flooring to grad-accum 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. --- openwam/train/openwam_trainer.py | 5 ++--- openwam/train/utils/checkpointing.py | 15 ++++++++------- tests/test_checkpointing.py | 19 ++++++++++++------- 3 files changed, 22 insertions(+), 17 deletions(-) diff --git a/openwam/train/openwam_trainer.py b/openwam/train/openwam_trainer.py index 2f8b3a55..66803215 100644 --- a/openwam/train/openwam_trainer.py +++ b/openwam/train/openwam_trainer.py @@ -536,9 +536,8 @@ def resume_if_configured(self, resume_state_dir: str, dataloader, grad_accum: in logger.info("[resume] loading Accelerate state from %s", resume_state_dir) meta = load_full_state(self.accelerator, resume_state_dir) global_step = int(meta.get("global_step", 0)) - # Align global_step to the grad_accum boundary skip was floored to, then derive - # opt_step from it — otherwise floored-off batches re-train and per-step seeds - # (keyed on global_step) drift. No-op at grad_accum=1. + # compute_resume_position keeps global_step exact (mid-accumulation included); + # derive opt_step from it so the LR schedule resumes where it actually was. start_epoch, skip, global_step = compute_resume_position(global_step, len(dataloader), grad_accum) opt_step = global_step // grad_accum if is_main: diff --git a/openwam/train/utils/checkpointing.py b/openwam/train/utils/checkpointing.py index 45827838..dbc895c3 100644 --- a/openwam/train/utils/checkpointing.py +++ b/openwam/train/utils/checkpointing.py @@ -320,17 +320,18 @@ def find_latest_accel_state(run_dir: str) -> str | None: def compute_resume_position(global_step: int, batches_per_epoch: int, grad_accum: int) -> tuple[int, int, int]: """Map a resumed ``global_step`` to ``(start_epoch, skip_first_batches, aligned_global_step)``. - ``skip`` is floored to a grad_accum boundary so the first optimizer step after - resume sees a full accumulation cycle; ``aligned_global_step`` pulls ``global_step`` - back to that same boundary so the floored-off batches are not re-trained and the - per-step seed (keyed on global_step) stays matched. No-op at grad_accum=1. + Resume continues from the exact checkpoint position: ``skip`` is the number of + batches already consumed in the epoch and ``aligned_global_step`` equals + ``global_step``. Accelerate's ``load_state`` restores the gradient-accumulation + counter, so a mid-accumulation resume keeps accumulating from where it left off; + flooring ``skip`` to a grad_accum boundary would re-feed already-consumed batches + and shift ``aligned_global_step`` (and the per-step seed keyed on it). ``grad_accum`` + is accepted for signature stability and is a no-op. """ batches_per_epoch = max(batches_per_epoch, 1) start_epoch = global_step // batches_per_epoch skip = global_step % batches_per_epoch - if grad_accum > 1 and skip % grad_accum != 0: - skip = (skip // grad_accum) * grad_accum - aligned_global_step = start_epoch * batches_per_epoch + skip + aligned_global_step = global_step return start_epoch, skip, aligned_global_step diff --git a/tests/test_checkpointing.py b/tests/test_checkpointing.py index e699d446..226afc4b 100644 --- a/tests/test_checkpointing.py +++ b/tests/test_checkpointing.py @@ -68,7 +68,7 @@ def test_find_latest_accel_state_none_when_empty(tmp_path): assert find_latest_accel_state(str(tmp_path)) is None -# --- compute_resume_position (grad_accum alignment / off-by fix) --- +# --- compute_resume_position (exact resume position / grad_accum is a no-op) --- def test_resume_position_grad_accum_1_is_identity(): @@ -76,14 +76,19 @@ def test_resume_position_grad_accum_1_is_identity(): assert compute_resume_position(25, 10, 1) == (2, 5, 25) -def test_resume_position_floors_skip_and_pulls_back_global_step(): - # gs=10, bpe=100, grad_accum=4: skip 10 -> 8, aligned 10 -> 8 (no re-train, step matched). - assert compute_resume_position(10, 100, 4) == (0, 8, 8) +def test_resume_position_keeps_exact_step_mid_accumulation(): + # gs=10, bpe=100, grad_accum=4: skip is the exact consumed count, aligned == gs. + assert compute_resume_position(10, 100, 4) == (0, 10, 10) -def test_resume_position_alignment_across_epoch(): - # gs=16, bpe=10, grad_accum=4: start=1, skip 6 -> 4, aligned = 1*10 + 4 = 14. - assert compute_resume_position(16, 10, 4) == (1, 4, 14) +def test_resume_position_exact_across_epoch(): + # gs=16, bpe=10, grad_accum=4: start=1, skip 6 (batches 10-15 consumed), aligned = 16. + assert compute_resume_position(16, 10, 4) == (1, 6, 16) + + +def test_resume_position_mid_accumulation_precise(): + # gs=12, bpe=10, grad_accum=4: epoch-1 batches 0-1 consumed, resume at 12 (not floored to 10). + assert compute_resume_position(12, 10, 4) == (1, 2, 12) def test_resume_position_already_aligned_unchanged(): From 05206b9ea990fb41ae86086c0b2a22866200dfa5 Mon Sep 17 00:00:00 2001 From: "LongZhenren(Zhibo Zhang)" Date: Thu, 8 Oct 2026 11:15:29 +0800 Subject: [PATCH 2/4] fix(train): degrade mid-cycle resume to the last sync boundary 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. --- openwam/train/openwam_trainer.py | 6 ++-- openwam/train/utils/checkpointing.py | 44 ++++++++++++++++++++++------ tests/test_checkpointing.py | 38 ++++++++++++++++++------ 3 files changed, 68 insertions(+), 20 deletions(-) diff --git a/openwam/train/openwam_trainer.py b/openwam/train/openwam_trainer.py index 66803215..15783a17 100644 --- a/openwam/train/openwam_trainer.py +++ b/openwam/train/openwam_trainer.py @@ -536,8 +536,10 @@ def resume_if_configured(self, resume_state_dir: str, dataloader, grad_accum: in logger.info("[resume] loading Accelerate state from %s", resume_state_dir) meta = load_full_state(self.accelerator, resume_state_dir) global_step = int(meta.get("global_step", 0)) - # compute_resume_position keeps global_step exact (mid-accumulation included); - # derive opt_step from it so the LR schedule resumes where it actually was. + # compute_resume_position floors mid-accumulation checkpoints back to the last + # sync boundary (the full state does not persist pending micro-batch gradients, + # so an exact mid-cycle resume would corrupt the next optimizer step); aligned + # boundary checkpoints resume exact. opt_step derives from the aligned step. start_epoch, skip, global_step = compute_resume_position(global_step, len(dataloader), grad_accum) opt_step = global_step // grad_accum if is_main: diff --git a/openwam/train/utils/checkpointing.py b/openwam/train/utils/checkpointing.py index dbc895c3..3de3fc86 100644 --- a/openwam/train/utils/checkpointing.py +++ b/openwam/train/utils/checkpointing.py @@ -320,18 +320,44 @@ def find_latest_accel_state(run_dir: str) -> str | None: def compute_resume_position(global_step: int, batches_per_epoch: int, grad_accum: int) -> tuple[int, int, int]: """Map a resumed ``global_step`` to ``(start_epoch, skip_first_batches, aligned_global_step)``. - Resume continues from the exact checkpoint position: ``skip`` is the number of - batches already consumed in the epoch and ``aligned_global_step`` equals - ``global_step``. Accelerate's ``load_state`` restores the gradient-accumulation - counter, so a mid-accumulation resume keeps accumulating from where it left off; - flooring ``skip`` to a grad_accum boundary would re-feed already-consumed batches - and shift ``aligned_global_step`` (and the per-step seed keyed on it). ``grad_accum`` - is accepted for signature stability and is a no-op. + Resumable checkpoints are only exact at real optimizer/sync boundaries. The + full-state snapshot persists the model, optimizer, and the Accelerate + step/accumulation counter, but *not* the pending micro-batch gradients of an + in-flight accumulation cycle: the training loop saves after every backward + whenever ``global_step % save_steps == 0``, 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 would see only the tail micro-batch (the + first parameter drifts in the opposite direction on a deterministic model). + Mid-cycle resumes are therefore degraded to the last global sync boundary: + ``aligned_global_step`` floors to the preceding ``grad_accum`` boundary and + ``(start_epoch, skip)`` are re-derived from it, and a warning is logged. + Boundary-aligned checkpoints (``global_step % grad_accum == 0``) resume exactly. + + This assumes the loop syncs on ``global_step % grad_accum == 0`` with no + epoch-end flush — true for ``OpenWAMTrainer``, where epochs roll over without + forcing an optimizer step, so ``global_step`` alone determines the sync + phase. If an epoch-end flush is ever added, this check must account for it. + ``grad_accum=1`` is unchanged (every step is a boundary). """ batches_per_epoch = max(batches_per_epoch, 1) - start_epoch = global_step // batches_per_epoch - skip = global_step % batches_per_epoch aligned_global_step = global_step + if grad_accum > 1 and aligned_global_step % grad_accum != 0: + aligned_global_step = (global_step // grad_accum) * grad_accum + start_epoch = aligned_global_step // batches_per_epoch + skip = aligned_global_step % batches_per_epoch + logger.warning( + "Checkpoint at global_step=%d is mid-accumulation; the full state does not " + "persist pending micro-batch gradients, so resuming exactly would corrupt " + "the next optimizer step. Degraded to the last sync boundary at " + "global_step=%d (skipping %d batches in the epoch).", + global_step, + aligned_global_step, + skip, + ) + return start_epoch, skip, aligned_global_step + start_epoch = aligned_global_step // batches_per_epoch + skip = aligned_global_step % batches_per_epoch return start_epoch, skip, aligned_global_step diff --git a/tests/test_checkpointing.py b/tests/test_checkpointing.py index 226afc4b..9d542683 100644 --- a/tests/test_checkpointing.py +++ b/tests/test_checkpointing.py @@ -68,34 +68,54 @@ def test_find_latest_accel_state_none_when_empty(tmp_path): assert find_latest_accel_state(str(tmp_path)) is None -# --- compute_resume_position (exact resume position / grad_accum is a no-op) --- +# --- compute_resume_position (boundary-aligned resume; mid-cycle degrades) --- def test_resume_position_grad_accum_1_is_identity(): - # grad_accum=1: aligned == global_step always (zero regression vs. pre-fix behaviour). + # grad_accum=1: every step is a sync boundary, aligned == global_step always. assert compute_resume_position(25, 10, 1) == (2, 5, 25) -def test_resume_position_keeps_exact_step_mid_accumulation(): - # gs=10, bpe=100, grad_accum=4: skip is the exact consumed count, aligned == gs. - assert compute_resume_position(10, 100, 4) == (0, 10, 10) +def test_resume_position_mid_cycle_degrades_to_last_sync_boundary(): + # gs=10, bpe=100, grad_accum=4: 10 % 4 == 2 (mid-cycle). The full state does not + # persist pending micro-batch gradients, so resume degrades to the last global + # sync boundary (gs=8) instead of resuming exactly. + assert compute_resume_position(10, 100, 4) == (0, 8, 8) -def test_resume_position_exact_across_epoch(): - # gs=16, bpe=10, grad_accum=4: start=1, skip 6 (batches 10-15 consumed), aligned = 16. +def test_resume_position_boundary_across_epoch_is_exact(): + # gs=16, bpe=10, grad_accum=4: 16 % 4 == 0 (real sync boundary), start=1, + # skip 6 (batches 10-15 consumed), resumes exactly at 16. assert compute_resume_position(16, 10, 4) == (1, 6, 16) -def test_resume_position_mid_accumulation_precise(): - # gs=12, bpe=10, grad_accum=4: epoch-1 batches 0-1 consumed, resume at 12 (not floored to 10). +def test_resume_position_global_boundary_inside_epoch_is_not_floored_by_skip(): + # gs=12, bpe=10, grad_accum=4: 12 % 4 == 0 — a global sync boundary even though + # the in-epoch skip (2) is not a multiple of grad_accum. The check must use the + # global step, so this resumes exactly at 12. assert compute_resume_position(12, 10, 4) == (1, 2, 12) +def test_resume_position_mid_cycle_across_epoch_degrades_globally(): + # gs=13, bpe=10, grad_accum=4: 13 % 4 == 1 (mid-cycle) -> degrade to the last + # global boundary 12, re-deriving epoch/skip from it (not flooring in-epoch + # skip 3 -> 0, which would lose the boundary at 12). + assert compute_resume_position(13, 10, 4) == (1, 2, 12) + + def test_resume_position_already_aligned_unchanged(): # gs=12, bpe=100, grad_accum=4: skip 12 already a multiple of 4 -> unchanged. assert compute_resume_position(12, 100, 4) == (0, 12, 12) +def test_resume_position_mid_cycle_logs_degrade_warning(caplog): + import logging + + with caplog.at_level(logging.WARNING, logger="openwam.train.utils.checkpointing"): + compute_resume_position(10, 100, 4) + assert any("mid-accumulation" in r.message and "Degraded" in r.message for r in caplog.records) + + # --- finalize_keep_weights_only --- From 099c0f08f18b0461c8ea90dabf7be00782c714dc Mon Sep 17 00:00:00 2001 From: "LongZhenren(Zhibo Zhang)" Date: Thu, 8 Oct 2026 12:12:53 +0800 Subject: [PATCH 3/4] fix(train): reject mid-cycle checkpoints on resume MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- openwam/train/openwam_trainer.py | 8 ++-- openwam/train/utils/checkpointing.py | 64 +++++++++++++++------------- tests/test_checkpointing.py | 62 +++++++++++++++------------ 3 files changed, 72 insertions(+), 62 deletions(-) diff --git a/openwam/train/openwam_trainer.py b/openwam/train/openwam_trainer.py index 15783a17..6bac98ed 100644 --- a/openwam/train/openwam_trainer.py +++ b/openwam/train/openwam_trainer.py @@ -536,10 +536,10 @@ def resume_if_configured(self, resume_state_dir: str, dataloader, grad_accum: in logger.info("[resume] loading Accelerate state from %s", resume_state_dir) meta = load_full_state(self.accelerator, resume_state_dir) global_step = int(meta.get("global_step", 0)) - # compute_resume_position floors mid-accumulation checkpoints back to the last - # sync boundary (the full state does not persist pending micro-batch gradients, - # so an exact mid-cycle resume would corrupt the next optimizer step); aligned - # boundary checkpoints resume exact. opt_step derives from the aligned step. + # compute_resume_position rejects mid-accumulation checkpoints (the full state + # does not persist pending micro-batch gradients, and no rewind can restore a + # consistent accumulation phase once load_state has set the Accelerator step); + # sync-boundary checkpoints resume exactly. opt_step derives from that step. start_epoch, skip, global_step = compute_resume_position(global_step, len(dataloader), grad_accum) opt_step = global_step // grad_accum if is_main: diff --git a/openwam/train/utils/checkpointing.py b/openwam/train/utils/checkpointing.py index 3de3fc86..5f775e20 100644 --- a/openwam/train/utils/checkpointing.py +++ b/openwam/train/utils/checkpointing.py @@ -320,45 +320,49 @@ def find_latest_accel_state(run_dir: str) -> str | None: def compute_resume_position(global_step: int, batches_per_epoch: int, grad_accum: int) -> tuple[int, int, int]: """Map a resumed ``global_step`` to ``(start_epoch, skip_first_batches, aligned_global_step)``. - Resumable checkpoints are only exact at real optimizer/sync boundaries. The + Only checkpoints taken at a real optimizer/sync boundary are resumable. The full-state snapshot persists the model, optimizer, and the Accelerate step/accumulation counter, but *not* the pending micro-batch gradients of an in-flight accumulation cycle: the training loop saves after every backward whenever ``global_step % save_steps == 0``, 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 would see only the tail micro-batch (the - first parameter drifts in the opposite direction on a deterministic model). - Mid-cycle resumes are therefore degraded to the last global sync boundary: - ``aligned_global_step`` floors to the preceding ``grad_accum`` boundary and - ``(start_epoch, skip)`` are re-derived from it, and a warning is logged. - Boundary-aligned checkpoints (``global_step % grad_accum == 0``) resume exactly. - - This assumes the loop syncs on ``global_step % grad_accum == 0`` with no - epoch-end flush — true for ``OpenWAMTrainer``, where epochs roll over without - forcing an optimizer step, so ``global_step`` alone determines the sync - phase. If an epoch-end flush is ever added, this check must account for it. - ``grad_accum=1`` is unchanged (every step is a boundary). + applied, so the next optimizer step would see only the tail micro-batch. + + A checkpoint at ``global_step`` is a sync boundary iff the 1-indexed position + within its epoch hits the accumulation quota or the epoch end — the prepared + dataloader runs with Accelerate's default ``sync_with_dataloader=True``, so + the epoch end flushes and resets the accumulation counter even when + ``batches_per_epoch % grad_accum != 0``:: + + local = ((global_step - 1) % batches_per_epoch) + 1 + is_sync_boundary = local % grad_accum == 0 or local == batches_per_epoch + + With ``batches_per_epoch=10, grad_accum=4`` the sync boundaries are exactly + {4, 8, 10, 14, 18, 20, ...} (measured on real ZeRO-1/ZeRO-2 runs). + + Mid-cycle checkpoints raise ``ValueError``: degrading them to an earlier + boundary cannot restore a consistent accumulation phase (the loaded + ``Accelerator.step`` still counts the consumed micro-batches), so the only + safe resume is from the nearest boundary checkpoint (checkpoint retention + keeps several). ``grad_accum=1`` makes every step a boundary. + + Raises: + ValueError: If the checkpoint was saved mid-accumulation-cycle. """ batches_per_epoch = max(batches_per_epoch, 1) - aligned_global_step = global_step - if grad_accum > 1 and aligned_global_step % grad_accum != 0: - aligned_global_step = (global_step // grad_accum) * grad_accum - start_epoch = aligned_global_step // batches_per_epoch - skip = aligned_global_step % batches_per_epoch - logger.warning( - "Checkpoint at global_step=%d is mid-accumulation; the full state does not " - "persist pending micro-batch gradients, so resuming exactly would corrupt " - "the next optimizer step. Degraded to the last sync boundary at " - "global_step=%d (skipping %d batches in the epoch).", - global_step, - aligned_global_step, - skip, + local_step = ((global_step - 1) % batches_per_epoch) + 1 + if grad_accum > 1 and local_step % grad_accum != 0 and local_step != batches_per_epoch: + raise ValueError( + f"Checkpoint at global_step={global_step} was saved mid-accumulation-cycle " + f"(position {local_step}/{batches_per_epoch} in the epoch, sync boundaries are " + f"every {grad_accum} batches or at the epoch end). Pending micro-batch gradients " + f"are not persisted in the full state, so this checkpoint cannot be resumed " + f"exactly. Use the nearest checkpoint at a sync boundary instead." ) - return start_epoch, skip, aligned_global_step - start_epoch = aligned_global_step // batches_per_epoch - skip = aligned_global_step % batches_per_epoch - return start_epoch, skip, aligned_global_step + start_epoch = global_step // batches_per_epoch + skip = global_step % batches_per_epoch + return start_epoch, skip, global_step # --- Retention / finalize (prune) --- diff --git a/tests/test_checkpointing.py b/tests/test_checkpointing.py index 9d542683..e48fef3e 100644 --- a/tests/test_checkpointing.py +++ b/tests/test_checkpointing.py @@ -68,7 +68,11 @@ def test_find_latest_accel_state_none_when_empty(tmp_path): assert find_latest_accel_state(str(tmp_path)) is None -# --- compute_resume_position (boundary-aligned resume; mid-cycle degrades) --- +# --- compute_resume_position (real sync boundaries; mid-cycle rejects) --- + +# Measured sync boundaries for bpe=10, ga=4 on real ZeRO-1/ZeRO-2 runs: +# updates fire at global steps {4, 8, 10, 14, 18, 20, ...} — every grad_accum +# batches, plus an epoch-end flush when bpe % grad_accum != 0. def test_resume_position_grad_accum_1_is_identity(): @@ -76,44 +80,46 @@ def test_resume_position_grad_accum_1_is_identity(): assert compute_resume_position(25, 10, 1) == (2, 5, 25) -def test_resume_position_mid_cycle_degrades_to_last_sync_boundary(): - # gs=10, bpe=100, grad_accum=4: 10 % 4 == 2 (mid-cycle). The full state does not - # persist pending micro-batch gradients, so resume degrades to the last global - # sync boundary (gs=8) instead of resuming exactly. - assert compute_resume_position(10, 100, 4) == (0, 8, 8) +def test_resume_position_accumulation_quota_boundary_is_exact(): + # gs=8, bpe=10, grad_accum=4: local position 8 hits the accumulation quota. + assert compute_resume_position(8, 10, 4) == (0, 8, 8) -def test_resume_position_boundary_across_epoch_is_exact(): - # gs=16, bpe=10, grad_accum=4: 16 % 4 == 0 (real sync boundary), start=1, - # skip 6 (batches 10-15 consumed), resumes exactly at 16. - assert compute_resume_position(16, 10, 4) == (1, 6, 16) +def test_resume_position_epoch_end_flush_boundary_is_exact(): + # gs=10, bpe=10, grad_accum=4: local position 10 == bpe. sync_with_dataloader + # flushes and resets the accumulation counter at the epoch end, so gs=10 is a + # real sync boundary even though 10 % 4 != 0. + assert compute_resume_position(10, 10, 4) == (1, 0, 10) -def test_resume_position_global_boundary_inside_epoch_is_not_floored_by_skip(): - # gs=12, bpe=10, grad_accum=4: 12 % 4 == 0 — a global sync boundary even though - # the in-epoch skip (2) is not a multiple of grad_accum. The check must use the - # global step, so this resumes exactly at 12. - assert compute_resume_position(12, 10, 4) == (1, 2, 12) +def test_resume_position_boundary_across_epoch_is_exact(): + # gs=14, bpe=10, grad_accum=4: local position 4 in epoch 1 (quota boundary). + assert compute_resume_position(14, 10, 4) == (1, 4, 14) + +def test_resume_position_epoch_end_boundary_in_later_epoch_is_exact(): + # gs=20, bpe=10, grad_accum=4: local position 10 == bpe in epoch 1. + assert compute_resume_position(20, 10, 4) == (2, 0, 20) -def test_resume_position_mid_cycle_across_epoch_degrades_globally(): - # gs=13, bpe=10, grad_accum=4: 13 % 4 == 1 (mid-cycle) -> degrade to the last - # global boundary 12, re-deriving epoch/skip from it (not flooring in-epoch - # skip 3 -> 0, which would lose the boundary at 12). - assert compute_resume_position(13, 10, 4) == (1, 2, 12) +@pytest.mark.parametrize("mid_cycle_step,epoch,skip", [(12, 1, 2), (16, 1, 6)]) +def test_resume_position_mid_cycle_checkpoint_rejects(mid_cycle_step, epoch, skip): + # gs=12 (local 2) and gs=16 (local 6) are mid-cycle: %4 != 0 and not the epoch + # end. The full state does not persist pending micro-batch gradients, and no + # rewind can restore a consistent accumulation phase once load_state has set + # the Accelerator step, so these must reject instead of resume. + with pytest.raises(ValueError, match="mid-accumulation-cycle"): + compute_resume_position(mid_cycle_step, 10, 4) -def test_resume_position_already_aligned_unchanged(): - # gs=12, bpe=100, grad_accum=4: skip 12 already a multiple of 4 -> unchanged. - assert compute_resume_position(12, 100, 4) == (0, 12, 12) +def test_resume_position_reject_message_names_boundary_advice(): + with pytest.raises(ValueError, match="sync boundary"): + compute_resume_position(16, 10, 4) -def test_resume_position_mid_cycle_logs_degrade_warning(caplog): - import logging - with caplog.at_level(logging.WARNING, logger="openwam.train.utils.checkpointing"): - compute_resume_position(10, 100, 4) - assert any("mid-accumulation" in r.message and "Degraded" in r.message for r in caplog.records) +def test_resume_position_grad_accum_1_never_rejects(): + # grad_accum=1: every step is a sync boundary even at an epoch edge. + assert compute_resume_position(10, 10, 1) == (1, 0, 10) # --- finalize_keep_weights_only --- From 099af408e9a51ceb7c943b94f18c0ba2d0ef5816 Mon Sep 17 00:00:00 2001 From: "LongZhenren(Zhibo Zhang)" Date: Fri, 9 Oct 2026 21:10:14 +0800 Subject: [PATCH 4/4] fix(train): write resumable states only at sync boundaries and skip unsafe states on resume MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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. --- openwam/train/openwam_trainer.py | 63 ++++++++++++++++++++++++-- openwam/train/utils/checkpointing.py | 44 ++++++++++++++++-- tests/test_checkpointing.py | 68 ++++++++++++++++++++++++++++ 3 files changed, 167 insertions(+), 8 deletions(-) diff --git a/openwam/train/openwam_trainer.py b/openwam/train/openwam_trainer.py index 6bac98ed..3bfb4511 100644 --- a/openwam/train/openwam_trainer.py +++ b/openwam/train/openwam_trainer.py @@ -22,6 +22,7 @@ """ import itertools +import json import logging import math import os @@ -32,7 +33,8 @@ from openwam.train.utils.checkpointing import ( compute_resume_position, finalize_keep_weights_only, - find_latest_accel_state, + find_accel_state_candidates, + is_sync_boundary, load_full_state, manage_checkpoints, save_config, @@ -358,11 +360,22 @@ def train(self, num_epochs: int = None, max_steps: int = None): ) # save_steps: write the weights line (+ the resumable full state when - # save_full_states_for_resume=true), then prune in lockstep. + # save_full_states_for_resume=true), then prune in lockstep. The full + # state is only written at a real sync boundary: a mid-cycle snapshot + # would persist no pending micro-batch gradients, and resuming it is + # rejected by compute_resume_position — so deferring here keeps every + # retained accel_state usable. 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) + else: + logger.info( + "[save] step %d is mid-accumulation; deferring the resumable " + "full state to the next sync boundary", + global_step, + ) if is_main: manage_checkpoints(output_path, keep_last_k) @@ -448,7 +461,11 @@ def setup_output_dir(self, debug: bool, resume_path: str | None) -> tuple[str, s base_output_path = getattr(self.cfg.training, "output_path", "./models") is_main = self.accelerator.is_main_process - resume_state_dir = find_latest_accel_state(resume_path) if resume_path else None + resume_state_candidates = find_accel_state_candidates(resume_path) if resume_path else [] + # Newest first; resume_if_configured walks this list past mid-cycle + # snapshots (written by older versions) to the nearest usable boundary. + self._resume_state_candidates = resume_state_candidates + resume_state_dir = resume_state_candidates[0][1] if resume_state_candidates else None if resume_path and resume_state_dir is None: raise FileNotFoundError( f"resume_ckpt_path={resume_path} has no usable accel_state_step_*; full states " @@ -534,7 +551,43 @@ def resume_if_configured(self, resume_state_dir: str, dataloader, grad_accum: in is_main = self.accelerator.is_main_process if is_main: logger.info("[resume] loading Accelerate state from %s", resume_state_dir) - meta = load_full_state(self.accelerator, resume_state_dir) + # Walk the candidate states newest-first and take the newest one that sits + # on a real sync boundary: compute_resume_position rejects mid-cycle + # snapshots (they persist no pending micro-batch gradients), which covers + # states written by older versions mid-accumulation-cycle. Only the chosen + # state is loaded, so a rejected candidate costs a metadata read, nothing more. + candidates = getattr(self, "_resume_state_candidates", None) or ( + [(None, resume_state_dir)] if resume_state_dir else [] + ) + chosen_dir: str | None = None + rejected: list[str] = [] + for _, state_dir in candidates: + meta_path = os.path.join(state_dir, "trainer_state.json") + try: + with open(meta_path) as f: + cand_step = int(json.load(f).get("global_step", 0)) + except (OSError, ValueError): + rejected.append(f"{os.path.basename(state_dir)} (unreadable trainer_state.json)") + continue + try: + compute_resume_position(cand_step, len(dataloader), grad_accum) + except ValueError as exc: + rejected.append(f"{os.path.basename(state_dir)} ({exc})") + continue + chosen_dir = state_dir + break + if chosen_dir is None: + raise ValueError( + f"No resumable sync-boundary state in {resume_state_dir}: " + + ("; ".join(rejected) if rejected else "no candidates") + ) + if is_main and chosen_dir != resume_state_dir: + logger.warning( + "[resume] newest state %s is mid-accumulation; falling back to %s", + os.path.basename(resume_state_dir), + os.path.basename(chosen_dir), + ) + meta = load_full_state(self.accelerator, chosen_dir) global_step = int(meta.get("global_step", 0)) # compute_resume_position rejects mid-accumulation checkpoints (the full state # does not persist pending micro-batch gradients, and no rewind can restore a diff --git a/openwam/train/utils/checkpointing.py b/openwam/train/utils/checkpointing.py index 5f775e20..a2eafe38 100644 --- a/openwam/train/utils/checkpointing.py +++ b/openwam/train/utils/checkpointing.py @@ -314,9 +314,47 @@ def find_latest_accel_state(run_dir: str) -> str | None: return best[1] if best else None +def find_accel_state_candidates(run_dir: str) -> list[tuple[int, str]]: + """Return all *finished* ``accel_state_step_N/`` dirs in *run_dir*, newest first. + + Same completion rule as ``find_latest_accel_state`` (``trainer_state.json`` + marker), but returns every finished candidate so the resume path can fall + back past mid-cycle snapshots that ``compute_resume_position`` rejects + (states written by older versions mid-accumulation-cycle). + """ + if not run_dir or not os.path.isdir(run_dir): + return [] + candidates: list[tuple[int, str]] = [] + for name in os.listdir(run_dir): + if not name.startswith("accel_state_step_"): + continue + state_dir = os.path.join(run_dir, name) + if not os.path.isfile(os.path.join(state_dir, "trainer_state.json")): + continue + candidates.append((step_num(name, "accel_state_step_"), state_dir)) + candidates.sort(key=lambda item: item[0], reverse=True) + return candidates + + # --- Resume position (pure math) --- +def is_sync_boundary(global_step: int, batches_per_epoch: int, grad_accum: int) -> bool: + """Whether ``global_step`` sits on a real optimizer/sync boundary. + + A step is a boundary iff the 1-indexed position within its epoch hits the + accumulation quota or the epoch end — the prepared dataloader runs with + Accelerate's default ``sync_with_dataloader=True``, so the epoch end flushes + and resets the accumulation counter even when + ``batches_per_epoch % grad_accum != 0``. With ``batches_per_epoch=10, + grad_accum=4`` the boundaries are exactly {4, 8, 10, 14, 18, 20, ...} + (measured on real ZeRO-1/ZeRO-2 runs). + """ + batches_per_epoch = max(batches_per_epoch, 1) + local_step = ((global_step - 1) % batches_per_epoch) + 1 + return local_step % grad_accum == 0 or local_step == batches_per_epoch + + def compute_resume_position(global_step: int, batches_per_epoch: int, grad_accum: int) -> tuple[int, int, int]: """Map a resumed ``global_step`` to ``(start_epoch, skip_first_batches, aligned_global_step)``. @@ -350,9 +388,9 @@ def compute_resume_position(global_step: int, batches_per_epoch: int, grad_accum Raises: ValueError: If the checkpoint was saved mid-accumulation-cycle. """ - batches_per_epoch = max(batches_per_epoch, 1) - local_step = ((global_step - 1) % batches_per_epoch) + 1 - if grad_accum > 1 and local_step % grad_accum != 0 and local_step != batches_per_epoch: + if not is_sync_boundary(global_step, batches_per_epoch, grad_accum): + batches_per_epoch = max(batches_per_epoch, 1) + local_step = ((global_step - 1) % batches_per_epoch) + 1 raise ValueError( f"Checkpoint at global_step={global_step} was saved mid-accumulation-cycle " f"(position {local_step}/{batches_per_epoch} in the epoch, sync boundaries are " diff --git a/tests/test_checkpointing.py b/tests/test_checkpointing.py index e48fef3e..51d3b475 100644 --- a/tests/test_checkpointing.py +++ b/tests/test_checkpointing.py @@ -11,8 +11,10 @@ from openwam.train.utils.checkpointing import ( compute_resume_position, finalize_keep_weights_only, + find_accel_state_candidates, find_latest_accel_state, find_latest_weights, + is_sync_boundary, load_full_state, manage_checkpoints, save_full_state, @@ -350,3 +352,69 @@ def test_setup_output_dir_verifies_stats_before_reusing_resume_run(tmp_path, mon assert output_path == str(run_dir) assert resume_state_dir == str(state_dir) assert calls == [(str(run_dir), trainer.dataset)] + + +# --- save -> prune -> discovery -> resume cross-link (bpe=10, ga=4) --- + + +def _make_resume_state(root: Path, step: int, global_step: int) -> Path: + d = root / f"accel_state_step_{step}" + d.mkdir(exist_ok=True) + (d / "trainer_state.json").write_text(f'{{"global_step": {global_step}}}') + return d + + +def test_sync_boundary_set_matches_measured_runs(): + # bpe=10, ga=4: measured sync boundaries on real ZeRO-1/ZeRO-2 runs. + boundaries = {4, 8, 10, 14, 18, 20} + for gs in range(1, 21): + assert is_sync_boundary(gs, 10, 4) == (gs in boundaries), gs + + +def test_save_prune_resume_chain_keeps_last_k_one_recovers_boundary_state(tmp_path): + # Save side defers the resumable full state past mid-cycle steps, so with + # keep_last_k=1 the weights line at the mid-cycle step (12) prunes the older + # weights (8) while the boundary state (8) is the only state left and survives. + _make_resume_state(tmp_path, 8, global_step=8) + (tmp_path / "checkpoint_step_8.safetensors").write_text("x") + (tmp_path / "checkpoint_step_12.safetensors").write_text("x") + manage_checkpoints(str(tmp_path), keep_last_k=1) + + assert (tmp_path / "checkpoint_step_12.safetensors").exists() + assert (tmp_path / "accel_state_step_8").is_dir() + + candidates = find_accel_state_candidates(str(tmp_path)) + chosen = next( + (state for step, state in candidates if _resume_allowed(step, 10, 4)), + None, + ) + assert chosen is not None and chosen.endswith("accel_state_step_8") + start_epoch, skip, aligned = compute_resume_position(8, 10, 4) + assert (start_epoch, skip, aligned) == (0, 8, 8) + + +def test_legacy_mid_cycle_state_falls_back_to_boundary_state(tmp_path): + # States written by older versions mid-accumulation-cycle sit in the run dir + # next to older boundary states: discovery returns both (newest first), the + # mid-cycle one is rejected by compute_resume_position, and the boundary + # state is resumed instead. + _make_resume_state(tmp_path, 8, global_step=8) + _make_resume_state(tmp_path, 12, global_step=12) + candidates = find_accel_state_candidates(str(tmp_path)) + assert [step for step, _ in candidates] == [12, 8] + + chosen = None + for step, state in candidates: + try: + compute_resume_position(step, 10, 4) + except ValueError: + continue + chosen = (step, state) + break + assert chosen is not None and chosen[0] == 8 + + +def _resume_allowed(step: int, bpe: int, ga: int) -> bool: + from openwam.train.utils.checkpointing import is_sync_boundary + + return is_sync_boundary(step, bpe, ga)