MCPcopy Create free account
hub / github.com/Audio-AGI/AudioSep / DataModule

Class DataModule

data/datamodules.py:8–82  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class DataModule(pl.LightningDataModule):
9 def __init__(
10 self,
11 train_dataset: object,
12 batch_size: int,
13 num_workers: int
14 ):
15 r"""Data module. To get one batch of data:
16
17 code-block:: python
18
19 data_module.setup()
20
21 for batch_data_dict in data_module.train_dataloader():
22 print(batch_data_dict.keys())
23 break
24
25 Args:
26 train_sampler: Sampler object
27 train_dataset: Dataset object
28 num_workers: int
29 distributed: bool
30 """
31 super().__init__()
32 self._train_dataset = train_dataset
33 self.num_workers = num_workers
34 self.batch_size = batch_size
35 self.collate_fn = collate_fn
36
37
38 def prepare_data(self):
39 # download, split, etc...
40 # only called on 1 GPU/TPU in distributed
41 pass
42
43 def setup(self, stage: Optional[str] = None) -> NoReturn:
44 r"""called on every device."""
45
46 # make assignments here (val/train/test split)
47 # called on every process in DDP
48
49 # SegmentSampler is used for selecting segments for training.
50 # On multiple devices, each SegmentSampler samples a part of mini-batch
51 # data.
52 self.train_dataset = self._train_dataset
53
54
55 def train_dataloader(self) -> torch.utils.data.DataLoader:
56 r"""Get train loader."""
57 train_loader = DataLoader(
58 dataset=self.train_dataset,
59 batch_size=self.batch_size,
60 collate_fn=self.collate_fn,
61 num_workers=self.num_workers,
62 pin_memory=True,
63 persistent_workers=False,
64 shuffle=True
65 )

Callers 1

get_data_moduleFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected