diff --git a/openwam/train/openwam_trainer.py b/openwam/train/openwam_trainer.py index 2f8b3a55..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,11 +551,48 @@ 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)) - # 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 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 45827838..a2eafe38 100644 --- a/openwam/train/utils/checkpointing.py +++ b/openwam/train/utils/checkpointing.py @@ -314,24 +314,93 @@ 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)``. - ``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. + 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. + + 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) + 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 " + 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." + ) 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 - return start_epoch, skip, aligned_global_step + return start_epoch, skip, global_step # --- Retention / finalize (prune) --- diff --git a/tests/test_checkpointing.py b/tests/test_checkpointing.py index e699d446..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, @@ -68,27 +70,58 @@ 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 (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(): - # 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_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_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_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_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_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) +@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_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 --- @@ -319,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)