refactor : replace iteration with consumed_samples

- Replace context.iteration with consumed_samples (global sample count)
- Add optimizer_step property derived from consumed_samples
- Checkpoint meta.json stores consumed_samples, drops iteration
- CLI --start_batch renamed to --start_samples (per-rank samples)
- Checkpoint dir naming: epoch_X_step_Y instead of epoch_X_iter_Y
- Metric log entries use step and consumed_samples fields
- Backward compat removed (old iteration checkpoints unsupported)
This commit is contained in:
2026-06-30 18:42:42 +08:00
parent 44579ea6dc
commit aabb0d83e9
10 changed files with 72 additions and 48 deletions
+9 -4
View File
@@ -25,7 +25,9 @@ def test_single_process():
scheduler.step()
checkpoint = Checkpoint(state_dict=model.state_dict(), epoch=3, iteration=30)
checkpoint = Checkpoint(
state_dict=model.state_dict(), epoch=3, consumed_samples=120
)
with tempfile.TemporaryDirectory() as tmpdir:
checkpoint.save(tmpdir)
@@ -33,7 +35,7 @@ def test_single_process():
loaded_checkpoint = Checkpoint.load(tmpdir)
assert loaded_checkpoint.epoch == 3
assert loaded_checkpoint.iteration == 30
assert loaded_checkpoint.consumed_samples == 120
def test_checkpoint_with_extra():
@@ -46,7 +48,10 @@ def test_checkpoint_with_extra():
"scheduler": {"last_epoch": 5},
}
checkpoint = Checkpoint(
state_dict=model.state_dict(), epoch=1, iteration=10, extra=extra
state_dict=model.state_dict(),
epoch=1,
consumed_samples=40,
extra=extra,
)
with tempfile.TemporaryDirectory() as tmpdir:
@@ -77,7 +82,7 @@ def simple_training():
checkpoint = Checkpoint(
state_dict=model.state_dict(),
epoch=2,
iteration=10,
consumed_samples=40,
)
rank = get_rank()