Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 62 additions & 8 deletions openwam/train/openwam_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
"""

import itertools
import json
import logging
import math
import os
Expand All @@ -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,
Expand Down Expand Up @@ -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)
Comment on lines 368 to +372

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.

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)

Expand Down Expand Up @@ -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 "
Expand Down Expand Up @@ -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)
Comment thread
wayrise marked this conversation as resolved.
opt_step = global_step // grad_accum
if is_main:
Expand Down
87 changes: 78 additions & 9 deletions openwam/train/utils/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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) ---
Expand Down
121 changes: 110 additions & 11 deletions tests/test_checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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 ---
Expand Down Expand Up @@ -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)
Loading