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
5 changes: 2 additions & 3 deletions openwam/train/openwam_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
15 changes: 8 additions & 7 deletions openwam/train/utils/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
19 changes: 12 additions & 7 deletions tests/test_checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,22 +68,27 @@ 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():
# grad_accum=1: aligned == global_step always (zero regression vs. pre-fix behaviour).
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():
Expand Down
Loading