-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtest_data_loader.py
More file actions
123 lines (91 loc) · 4.26 KB
/
Copy pathtest_data_loader.py
File metadata and controls
123 lines (91 loc) · 4.26 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
import pytest
import numpy as np
from unittest.mock import MagicMock, patch
from src.data_loader import DataLoader
from datasets import Dataset, DatasetDict
from PIL import Image
import torch
from torch.utils.data import DataLoader as TorchDataLoader
@pytest.fixture
def mock_hugging_face_dataset():
"""Creates a dataset with 2 rows for test and train data"""
train_image_1 = Image.new('RGB', (10, 10), color = 'red')
train_image_2 = Image.new('RGB', (10,10), color = 'blue')
test_image_1 = Image.new('RGB', (10, 10), color = 'red')
test_image_2 = Image.new('RGB', (10,10), color = 'blue')
train_dataset = Dataset.from_dict({
'image': [np.array(train_image_1), np.array(train_image_2)],
'label': [0,1]
})
test_dataset = Dataset.from_dict({
'image': [np.array(test_image_1), np.array(test_image_2)],
'label': [0,1]
})
return Dataset.from_dict({ 'train': train_dataset, 'test': test_dataset })
@patch('src.data_loader.load_from_disk')
def test_load_with_assets_in_folder(mock_load_from_disk, tmp_path):
with patch('src.data_loader.RAW_DATA_DIR', tmp_path):
loader = DataLoader()
fake_file = tmp_path / "dummy.arrow"
fake_file.parent.mkdir(parents = True, exist_ok = True)
fake_file.touch()
mock_dataset = MagicMock()
mock_load_from_disk.return_value = mock_dataset
result = loader.load()
expected_string_path = str(tmp_path)
mock_load_from_disk.assert_called_once_with(expected_string_path)
assert result == mock_dataset
def test_load_with_no_data_in_folder(tmp_path):
with patch('src.data_loader.RAW_DATA_DIR', tmp_path):
data_loader = DataLoader()
with pytest.raises(FileNotFoundError) as exec_info:
data_loader.load()
@patch('src.data_loader.load_from_disk')
def test_get_train_data(mock_load_from_disk, tmp_path, mock_hugging_face_dataset):
with patch('src.data_loader.RAW_DATA_DIR', tmp_path):
loader = DataLoader()
fake_file = tmp_path / "dummy.arrow"
fake_file.parent.mkdir(parents = True, exist_ok = True)
fake_file.touch()
mock_load_from_disk.return_value = mock_hugging_face_dataset
X_train, y_train = loader.get_train_data()
assert len(X_train) == 2
assert len(y_train) == 2
@patch('src.data_loader.load_from_disk')
def test_get_train_loader(mock_load_from_disk, tmp_path, mock_hugging_face_dataset):
with patch('src.data_loader.RAW_DATA_DIR', tmp_path):
expected_batch_size = 2
expected_workers = 2
data_loader = DataLoader(
batch_size = expected_batch_size,
num_workers = expected_workers
)
loader = DataLoader()
fake_file = tmp_path / "dummy.arrow"
fake_file.parent.mkdir(parents = True, exist_ok = True)
fake_file.touch()
mock_load_from_disk.return_value = mock_hugging_face_dataset
train_loader = data_loader.get_train_loader()
assert isinstance(train_loader, TorchDataLoader), "Returns a PyTorch DataLoader object"
assert train_loader.batch_size == expected_batch_size, 'Batch size mismatch'
assert train_loader.num_workers == expected_workers, 'Worker allocation mismatch'
assert train_loader.drop_last is False
@patch('src.data_loader.load_from_disk')
def test_get_test_loader(mock_load_from_disk, tmp_path, mock_hugging_face_dataset):
with patch('src.data_loader.RAW_DATA_DIR', tmp_path):
expected_batch_size = 2
expected_workers = 2
data_loader = DataLoader(
batch_size = expected_batch_size,
num_workers = expected_workers
)
loader = DataLoader()
fake_file = tmp_path / "dummy.arrow"
fake_file.parent.mkdir(parents = True, exist_ok = True)
fake_file.touch()
mock_load_from_disk.return_value = mock_hugging_face_dataset
test_loader = data_loader.get_test_loader()
assert isinstance(test_loader, TorchDataLoader), "Returns a PyTorch DataLoader object"
assert test_loader.batch_size == expected_batch_size, 'Batch size mismatch'
assert test_loader.num_workers == expected_workers, 'Worker allocation mismatch'
assert test_loader.drop_last is False