-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathChexpertDataset.py
More file actions
104 lines (86 loc) · 3.08 KB
/
Copy pathChexpertDataset.py
File metadata and controls
104 lines (86 loc) · 3.08 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
from torch.utils.data import Dataset
from utilFunc import readCSV
import argParser as ARG
from torchvision.transforms import v2
from PIL import Image
import torchvision.transforms.functional
import torch
TEST = "test"
TRAIN = "train"
VALIDATE = "validate"
POSSIBLE_SPLITS = [
TEST,
TRAIN,
VALIDATE,
]
if ARG.ON_HPC:
BASE_PATH = "/hpcwork/p0021834/workspace_patrick/datasets/CheXpert-v1.0_1024/"
else:
BASE_PATH = "/home/patrick/datasets/CheXpert-v1.0_1024/"
META_DICT = {
TRAIN:"train_visualCheXbert.csv",
#TRAIN:"train.csv", #for the classic training data
VALIDATE:"valid.csv",
TEST:"test_labels.csv",
}
META_LABEL_CUT = 5
META_LABEL_CUT_DICT = {
TRAIN:5,
VALIDATE:5,
TEST:1,
}
def shiftVisualChexpert(csvCont): # align visualCheXpert label with the normal chexpert training label order
labelCut = 5
for line in csvCont:
lastEl = line.pop()
line.insert(labelCut,lastEl)
print("SHIFTED VISUAL CHEXPERT TRAINING FILE")
class ClassificationChexpertDataset(Dataset):
def __init__(self, usedSplit, augementImg=False, imgSize=ARG.IMG_SIZE):
self.usedSplit = usedSplit
self.augementImg = augementImg
self.imgSize = imgSize
metaFile = BASE_PATH + META_DICT[self.usedSplit]
#head,*data = readCSV(metaFile)
csvCont = readCSV(metaFile)
if META_DICT[self.usedSplit] == "train_visualCheXbert.csv": shiftVisualChexpert(csvCont) #shift dataset
head,*data = csvCont
labelCut = META_LABEL_CUT_DICT[usedSplit]
self.labelNames = head[labelCut:]
getTotalPath = lambda relativePath: BASE_PATH + relativePath.replace("CheXpert-v1.0/","")
labelDict = { #set unknown to 0 => there is no unknow in in the visual chexpert label
"": 0,
"1.0": 1,
"0.0": 0,
"-1.0": 0,
}
#labelToInt = lambda lst: [int(float(x)) for x in lst]
labelToInt = lambda lst: [labelDict[x] for x in lst]
self.itemList = [(getTotalPath(x[0]),labelToInt(x[labelCut:])) for x in data]
self.augmentTransforms = v2.Compose([
v2.RandomResizedCrop(size=(self.imgSize,self.imgSize),scale=(0.5,1),antialias=True),
v2.RandomRotation(degrees=5),
v2.ColorJitter(brightness=0.3),
])
self.normalTransform = v2.Compose([
v2.Resize(size=(self.imgSize,self.imgSize),antialias=True)
])
def applyAugment(self,img):
if not self.augementImg:
return self.normalTransform(img)
img = self.augmentTransforms(img)
return img
def getLabelNames(self):
return self.labelNames
def __len__(self):
return len(self.itemList)
def __getitem__(self, idx):
imgPath,label = self.itemList[idx]
img = self.getImg(imgPath)
img = self.applyAugment(img)
label = torch.tensor(label,dtype=torch.float32)
return img,label
def getImg(self,imgPath):
img = Image.open(imgPath).convert("L")
img = torchvision.transforms.functional.to_tensor(img)
return img