-
-
Notifications
You must be signed in to change notification settings - Fork 257
Expand file tree
/
Copy pathtest_dataloaders.py
More file actions
66 lines (44 loc) · 2.07 KB
/
Copy pathtest_dataloaders.py
File metadata and controls
66 lines (44 loc) · 2.07 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
import pytest
import torchxrayvision as xrv
from skimage.io import imread, imsave
dataset_classes = [xrv.datasets.NIH_Dataset,
xrv.datasets.PC_Dataset,
xrv.datasets.NIH_Google_Dataset,
xrv.datasets.Openi_Dataset]
def test_dataloader_basic():
for dataset_class in dataset_classes:
dataset_class(imgpath=".")
def test_dataloader_merging():
datasets = []
for dataset_class in dataset_classes:
dataset = dataset_class(imgpath=".")
datasets.append(dataset)
for dataset in datasets:
xrv.datasets.relabel_dataset(xrv.datasets.default_pathologies, dataset)
dd = xrv.datasets.Merge_Dataset(datasets)
# test that we catch incorrect pathology alignment
def test_dataloader_merging_incorrect_alignment():
with pytest.raises(Exception) as excinfo:
d_nih = xrv.datasets.NIH_Dataset(imgpath=".")
d_pc = xrv.datasets.PC_Dataset(imgpath=".")
dd = xrv.datasets.Merge_Dataset([d_nih, d_pc])
assert "incorrect pathology alignment" in str(excinfo.value)
with pytest.raises(Exception) as excinfo:
d_nih = xrv.datasets.NIH_Dataset(imgpath=".")
d_pc = xrv.datasets.PC_Dataset(imgpath=".")
xrv.datasets.relabel_dataset(xrv.datasets.default_pathologies, d_nih)
xrv.datasets.relabel_dataset(xrv.datasets.default_pathologies[:-1], d_pc)
dd = xrv.datasets.Merge_Dataset([d_nih, d_pc])
assert "incorrect pathology alignment" in str(excinfo.value)
def test_resize():
for filename in ["16747_3_1.jpg", "covid-19-pneumonia-58-prior.jpg"]
img = imread(filename)
img = xrv.datasets.normalize(img, 255)
# Check that images are 2D arrays
if len(img.shape) > 2:
img = img[:, :, 0]
# Add color channel
img = img[None, :, :]
resize_ski = xrv.datasets.XRayResizer(100, engine="skimage")
resize_cv2 = xrv.datasets.XRayResizer(100, engine="cv2")
assert(np.allclose(resize_ski(img),resize_cv2(img)))