diff --git a/astrai/config/train_config.py b/astrai/config/train_config.py index f08e742..084d84a 100644 --- a/astrai/config/train_config.py +++ b/astrai/config/train_config.py @@ -48,6 +48,7 @@ class TrainConfig(BaseConfig): random_seed (int): Random seed. Defaults to 3407. num_workers (int): Number of workers for dataloader. Defaults to 0. prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None. + persistent_workers (bool): Keep DataLoader workers alive between epochs. Defaults to False. pin_memory (bool): Pin memory for dataloader. Defaults to False. collate_fn (Optional[Callable[[List[Any]], Any]]): Collate function for dataloader (e.g. dpo_collate_fn). Defaults to None. nprocs (int): Number of processes for distributed training. Defaults to 1. @@ -98,6 +99,7 @@ class TrainConfig(BaseConfig): random_seed: int = 3407 num_workers: int = 0 prefetch_factor: Optional[int] = None + persistent_workers: bool = False pin_memory: bool = False collate_fn: Optional[Callable[[List[Any]], Any]] = None diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index c4d4f98..99f65f5 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -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: diff --git a/scripts/tools/train.py b/scripts/tools/train.py index fd0402b..10cdbb3 100644 --- a/scripts/tools/train.py +++ b/scripts/tools/train.py @@ -251,6 +251,12 @@ _START_METHODS = ["spawn", "fork", "forkserver"] group="Data Loading", help="Pin memory.", ) +@opt( + "--persistent_workers/--no-persistent_workers", + default=True, + group="Data Loading", + help="Keep DataLoader workers alive between epochs.", +) @opt( "--window_size", type=int, @@ -606,6 +612,7 @@ def train( random_seed: int, num_workers: int, pin_memory: bool, + persistent_workers: bool, gradient_checkpointing: bool, window_size: int, stride: int, @@ -797,6 +804,7 @@ def train( random_seed=random_seed, num_workers=num_workers, pin_memory=pin_memory, + persistent_workers=persistent_workers, nprocs=nprocs, backend=backend, master_addr=master_addr,