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:
@@ -48,6 +48,7 @@ class TrainConfig(BaseConfig):
|
|||||||
random_seed (int): Random seed. Defaults to 3407.
|
random_seed (int): Random seed. Defaults to 3407.
|
||||||
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
num_workers (int): Number of workers for dataloader. Defaults to 0.
|
||||||
prefetch_factor (Optional[int]): Prefetch factor for dataloader. Defaults to None.
|
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.
|
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.
|
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.
|
nprocs (int): Number of processes for distributed training. Defaults to 1.
|
||||||
@@ -98,6 +99,7 @@ class TrainConfig(BaseConfig):
|
|||||||
random_seed: int = 3407
|
random_seed: int = 3407
|
||||||
num_workers: int = 0
|
num_workers: int = 0
|
||||||
prefetch_factor: Optional[int] = None
|
prefetch_factor: Optional[int] = None
|
||||||
|
persistent_workers: bool = False
|
||||||
pin_memory: bool = False
|
pin_memory: bool = False
|
||||||
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
collate_fn: Optional[Callable[[List[Any]], Any]] = None
|
||||||
|
|
||||||
|
|||||||
@@ -231,15 +231,22 @@ class TrainContextBuilder:
|
|||||||
seed=cfg.random_seed,
|
seed=cfg.random_seed,
|
||||||
shuffle=shuffle,
|
shuffle=shuffle,
|
||||||
)
|
)
|
||||||
return DataLoader(
|
loader_kwargs = dict(
|
||||||
dataset,
|
dataset=dataset,
|
||||||
batch_size=cfg.batch_per_device,
|
batch_size=cfg.batch_per_device,
|
||||||
sampler=sampler,
|
sampler=sampler,
|
||||||
num_workers=cfg.num_workers,
|
num_workers=cfg.num_workers,
|
||||||
pin_memory=cfg.pin_memory,
|
pin_memory=cfg.pin_memory,
|
||||||
prefetch_factor=cfg.prefetch_factor,
|
|
||||||
collate_fn=cfg.collate_fn,
|
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:
|
def _restore_optimizer_state(self, context: TrainContext) -> None:
|
||||||
if context.checkpoint and context.checkpoint.extra:
|
if context.checkpoint and context.checkpoint.extra:
|
||||||
|
|||||||
@@ -251,6 +251,12 @@ _START_METHODS = ["spawn", "fork", "forkserver"]
|
|||||||
group="Data Loading",
|
group="Data Loading",
|
||||||
help="Pin memory.",
|
help="Pin memory.",
|
||||||
)
|
)
|
||||||
|
@opt(
|
||||||
|
"--persistent_workers/--no-persistent_workers",
|
||||||
|
default=True,
|
||||||
|
group="Data Loading",
|
||||||
|
help="Keep DataLoader workers alive between epochs.",
|
||||||
|
)
|
||||||
@opt(
|
@opt(
|
||||||
"--window_size",
|
"--window_size",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -606,6 +612,7 @@ def train(
|
|||||||
random_seed: int,
|
random_seed: int,
|
||||||
num_workers: int,
|
num_workers: int,
|
||||||
pin_memory: bool,
|
pin_memory: bool,
|
||||||
|
persistent_workers: bool,
|
||||||
gradient_checkpointing: bool,
|
gradient_checkpointing: bool,
|
||||||
window_size: int,
|
window_size: int,
|
||||||
stride: int,
|
stride: int,
|
||||||
@@ -797,6 +804,7 @@ def train(
|
|||||||
random_seed=random_seed,
|
random_seed=random_seed,
|
||||||
num_workers=num_workers,
|
num_workers=num_workers,
|
||||||
pin_memory=pin_memory,
|
pin_memory=pin_memory,
|
||||||
|
persistent_workers=persistent_workers,
|
||||||
nprocs=nprocs,
|
nprocs=nprocs,
|
||||||
backend=backend,
|
backend=backend,
|
||||||
master_addr=master_addr,
|
master_addr=master_addr,
|
||||||
|
|||||||
Reference in New Issue
Block a user