-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathDataModule.py
More file actions
130 lines (116 loc) · 4.64 KB
/
Copy pathDataModule.py
File metadata and controls
130 lines (116 loc) · 4.64 KB
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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
from pytorch_lightning.core import datamodule
import torch
import pytorch_lightning as pl
import os
import shutil
from dataset import VOCDataset
class DataModule(pl.LightningDataModule):
def __init__(self,
root_dir,
batch_size,
spatial_resolution,
mean = [0.485, 0.456, 0.406],
std = [0.229, 0.224, 0.225],
num_workers=None,
pin_memory=False,
shuffle=True,
cache_dir=None,
cache_refresh=True,
):
super().__init__()
self.root_dir = root_dir
self.cache_dir = cache_dir
self.cache_refresh = cache_refresh
self.csv_file = {
'train': 'train.csv',
'val': 'test.csv',
'test': 'test.csv',
}
self.spatial_resolution = spatial_resolution
self.mean = mean
self.std = std
self.batch_size = batch_size
self.num_workers = os.cpu_count() - 1 if num_workers is None else num_workers
self.pin_memory = pin_memory
self.shuffle = shuffle
def prepare_data(self) -> None:
if self.cache_dir is not None:
if self.cache_refresh == True or os.path.exists(self.cache_dir) == False:
shutil.rmtree(self.cache_dir, ignore_errors=True)
os.mkdir(self.cache_dir)
def setup(self, stage = None):
if stage == "fit" or stage is None:
self.train_dataset = VOCDataset(
root_dir=self.root_dir,
csv_file=self.csv_file['train'],
cache_dir=os.path.join(self.cache_dir, 'train') if self.cache_dir is not None else None,
cache_refresh=self.cache_refresh,
spatial_resolution=self.spatial_resolution,
mean=self.mean,
std=self.std,
augmentation=True,
)
self.val_dataset = VOCDataset(
root_dir=self.root_dir,
csv_file=self.csv_file['val'],
cache_dir=os.path.join(self.cache_dir, 'val') if self.cache_dir is not None else None,
cache_refresh=self.cache_refresh,
spatial_resolution=self.spatial_resolution,
mean=self.mean,
std=self.std,
augmentation=False,
)
if stage == "test" or stage is None:
self.test_dataset = VOCDataset(
root_dir=self.root_dir,
csv_file=self.csv_file['test'],
cache_dir=os.path.join(self.cache_dir, 'test') if self.cache_dir is not None else None,
cache_refresh=self.cache_refresh,
spatial_resolution=self.spatial_resolution,
mean=self.mean,
std=self.std,
augmentation=False,
)
def train_dataloader(self):
return torch.utils.data.DataLoader(
self.train_dataset,
batch_size=self.batch_size,
shuffle=self.shuffle,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
drop_last=True,
)
def val_dataloader(self):
return torch.utils.data.DataLoader(
self.val_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
)
def test_dataloader(self):
return torch.utils.data.DataLoader(
self.test_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=self.num_workers,
pin_memory=self.pin_memory,
)
if __name__ == '__main__':
datamodule = DataModule(
root_dir="../VOC100examples",
cache_dir='cache',
cache_refresh=False,
spatial_resolution=[512, 512],
batch_size=1,
num_workers=0,
pin_memory=True,
)
datamodule.prepare_data()
datamodule.setup()
train_dl = datamodule.train_dataloader()
print('len', len(train_dl))
for data in train_dl:
break
for k in data.keys():
print(k, data[k].shape)