MCPcopy Create free account
hub / github.com/dcharatan/flowmap / DataModulePretrain

Class DataModulePretrain

flowmap/dataset/data_module_pretrain.py:39–84  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

37
38
39class DataModulePretrain(LightningDataModule):
40 def __init__(
41 self,
42 dataset_cfgs: list[DatasetCfg],
43 data_module_cfg: DataModulePretrainCfg,
44 frame_sampler_cfg: FrameSamplerCfg,
45 global_rank: int,
46 ) -> None:
47 super().__init__()
48 self.dataset_cfgs = dataset_cfgs
49 self.data_module_cfg = data_module_cfg
50 self.frame_sampler_cfg = frame_sampler_cfg
51 self.global_rank = global_rank
52
53 def get_persistent(self, loader_cfg: DataLoaderStageCfg) -> bool | None:
54 return None if loader_cfg.num_workers == 0 else loader_cfg.persistent_workers
55
56 def get_generator(self, loader_cfg: DataLoaderStageCfg) -> torch.Generator | None:
57 if loader_cfg.seed is None:
58 return None
59 generator = Generator()
60 generator.manual_seed(loader_cfg.seed + self.global_rank)
61 return generator
62
63 def train_dataloader(self):
64 dataset = get_dataset(self.dataset_cfgs, "train", self.frame_sampler_cfg)
65 return DataLoader(
66 dataset,
67 self.data_module_cfg.train.batch_size,
68 shuffle=not isinstance(dataset, IterableDataset),
69 num_workers=self.data_module_cfg.train.num_workers,
70 generator=self.get_generator(self.data_module_cfg.train),
71 worker_init_fn=worker_init_fn,
72 persistent_workers=self.get_persistent(self.data_module_cfg.train),
73 )
74
75 def val_dataloader(self):
76 dataset = get_dataset(self.dataset_cfgs, "val", self.frame_sampler_cfg)
77 return DataLoader(
78 ValidationWrapper(dataset, 1),
79 self.data_module_cfg.val.batch_size,
80 num_workers=self.data_module_cfg.val.num_workers,
81 generator=self.get_generator(self.data_module_cfg.val),
82 worker_init_fn=worker_init_fn,
83 persistent_workers=self.get_persistent(self.data_module_cfg.val),
84 )

Callers 1

pretrainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected