fix: correct epoch computation on resume to avoid redoing whole epoch

This commit is contained in:
2026-07-28 00:01:29 +08:00
parent 2e29ed45d3
commit 5ab18bec48
+16 -12
View File
@@ -101,18 +101,13 @@ class TrainContextBuilder:
if checkpoint.config: if checkpoint.config:
model_config = checkpoint.config model_config = checkpoint.config
if self._resume: if self._resume:
preloaded_epoch = checkpoint.epoch or cfg.start_epoch preloaded_epoch = checkpoint.epoch
if checkpoint.consumed_samples > 0: per_step = (
per_step = ( cfg.batch_per_device * get_world_size() * cfg.grad_accum_steps
cfg.batch_per_device )
* get_world_size() preloaded_consumed = (
* cfg.grad_accum_steps checkpoint.consumed_samples // per_step
) ) * per_step
preloaded_consumed = (
checkpoint.consumed_samples // per_step
) * per_step
else:
preloaded_consumed = cfg.start_samples * get_world_size()
preloaded_checkpoint = checkpoint preloaded_checkpoint = checkpoint
if not model_config and hasattr(cfg.model_fn(), "config"): if not model_config and hasattr(cfg.model_fn(), "config"):
@@ -162,6 +157,15 @@ class TrainContextBuilder:
) )
sampler_offset = context.consumed_samples // context.world_size sampler_offset = context.consumed_samples // context.world_size
if self._resume and sampler_offset > 0:
offset = context.world_size - 1
num_samples_per_replica = (
len(train_dataset) + offset
) // context.world_size
if num_samples_per_replica > 0:
context.epoch = sampler_offset // num_samples_per_replica
sampler = RDSampler( sampler = RDSampler(
data_source=train_dataset, data_source=train_dataset,
start_epoch=context.epoch, start_epoch=context.epoch,