feat: add persistent DataLoader workers

- Keep training workers alive between epochs when enabled.
- Avoid invalid prefetch settings for single-process loading.
This commit is contained in:
2026-08-09 11:38:50 +08:00
parent 3416f98c58
commit cf4f5ab9f6
3 changed files with 20 additions and 3 deletions
+10 -3
View File
@@ -231,15 +231,22 @@ class TrainContextBuilder:
seed=cfg.random_seed,
shuffle=shuffle,
)
return DataLoader(
dataset,
loader_kwargs = dict(
dataset=dataset,
batch_size=cfg.batch_per_device,
sampler=sampler,
num_workers=cfg.num_workers,
pin_memory=cfg.pin_memory,
prefetch_factor=cfg.prefetch_factor,
collate_fn=cfg.collate_fn,
)
# PyTorch rejects prefetch_factor/persistent_workers when workers=0.
if cfg.num_workers > 0:
loader_kwargs["persistent_workers"] = cfg.persistent_workers
if cfg.prefetch_factor is not None:
loader_kwargs["prefetch_factor"] = cfg.prefetch_factor
return DataLoader(
**loader_kwargs,
)
def _restore_optimizer_state(self, context: TrainContext) -> None:
if context.checkpoint and context.checkpoint.extra: