File size: 2,818 Bytes
59630ba | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | import torch
from utils.distributed_utils import rank_zero_print
from utils.print_utils import cyan
from .base_data_module import BaseDataModule
class ResumableDataModule(BaseDataModule):
"""
A resumable data module for `datasets.video.base_iterable_video_dataset.BaseAdvancedVideoDataset`.
Activated when `experiment.reload_dataloaders_every_n_epochs = 1`, `experiment.training.data.shuffle = False`, checkpointing & validation are epoch-based, and `dataset.subdataset_size` is set.
Data module will pass the current epoch to the dataset, so that the dataset can deterministically compute the subdataset corresponding to the current epoch.
"""
@property
def is_resumable(self) -> bool:
is_experiment_resumable = (
self.root_cfg.experiment.reload_dataloaders_every_n_epochs == 1
and not self.exp_cfg.training.data.shuffle
and self.root_cfg.experiment.training.checkpointing.every_n_epochs
is not None
and self.root_cfg.experiment.validation.val_every_n_epoch is not None
and self.root_cfg.experiment.training.max_steps == -1
and (
self.root_cfg.experiment.validation.val_every_n_epoch > 1
or self.root_cfg.experiment.validation.val_every_n_step == 1.0
) # this ensures that ckpts are saved not before validation (otherwise,it may lead to redundant ckpts & validation if resuming)
)
is_dataset_resumable = self.root_cfg.dataset.subdataset_size is not None
if is_experiment_resumable != is_dataset_resumable:
raise ValueError(
"To make a resumable run, set experiment.reload_dataloaders_every_n_epochs = 1, experiment.training.data.shuffle = False, checkpointing & validation are epoch-based, and dataset.subdataset_size is set."
)
return is_experiment_resumable and is_dataset_resumable
def _build_dataset(self, split: str) -> torch.utils.data.Dataset:
if split in ["training", "test", "validation"]:
is_resumable = self.is_resumable and split == "training"
dataset = self.compatible_datasets[self.root_cfg.dataset._name](
self.root_cfg.dataset,
split=split,
current_epoch=(self.trainer.current_epoch if is_resumable else None),
)
if is_resumable:
rank_zero_print(
cyan(
f"Resumable Training ({(dataset.cumulative_sizes[-1] / dataset.subdataset_size):.1f} subepochs / epoch)"
),
f"currently at subepoch {dataset.current_subepoch}",
)
return dataset
else:
raise NotImplementedError(f"split '{split}' is not implemented")
|