-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdataset.py
More file actions
116 lines (94 loc) · 3.36 KB
/
Copy pathdataset.py
File metadata and controls
116 lines (94 loc) · 3.36 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
import csv
import os
import torch
from torch.utils.data import Dataset, DataLoader
from torchvision.io import read_image
import torchvision.transforms as transforms
import numpy as np
import matplotlib.pyplot as plt
SPEED_SCALE = 400
class TrackManiaDataset(Dataset):
def __init__(
self,
data_dir,
annotations_file_name,
only_steer=False,
transform=None,
target_transform=None,
):
anno_file = open(os.path.join(data_dir, annotations_file_name))
self.img_labels = list(csv.DictReader(anno_file))
anno_file.close()
self.data_dir = data_dir
self.transform = transform
self.target_transform = target_transform
self.only_steer = only_steer
def __len__(self):
return len(self.img_labels)
def __getitem__(self, idx):
row = self.img_labels[idx]
img_path = os.path.join(self.data_dir, row["img_file"])
image = read_image(img_path)
speed = float(row["speed"]) / SPEED_SCALE
steering = float(row["steering"])
if self.transform:
image = self.transform(image)
if self.target_transform:
speed = self.target_transform(speed)
steering = self.target_transform(steering)
if self.only_steer:
return image, torch.from_numpy(np.array([steering])).float()
return image, torch.from_numpy(np.array([speed, steering])).float()
def view_data(data):
figure = plt.figure(figsize=(8, 8))
cols, rows = 3, 3
for i in range(1, cols * rows + 1):
sample_idx = torch.randint(len(data), size=(1,)).item()
img, (speed, steering) = data[sample_idx]
figure.add_subplot(rows, cols, i)
plt.title("Speed: %.3f\nSteering: %.3f" % (float(speed), float(steering)))
plt.axis("off")
plt.imshow(img.squeeze(), cmap="gray")
plt.show()
def view_dataloader(dataloader):
features, labels = next(iter(dataloader))
speed_labels = labels[:, 0]
steering_labels = labels[:, 1]
print(f"Feature batch shape: {features.size()}")
print(f"Speed batch shape: {speed_labels.size()}")
print(f"Steering batch shape: {steering_labels.size()}")
img = features[0].squeeze()
speed = speed_labels[0]
steering = steering_labels[0]
plt.imshow(img, cmap="gray")
plt.title("Speed: %.3f\nSteering: %.3f" % (float(speed), float(steering)))
plt.show()
def load_data():
training_data = TrackManiaDataset(
"data",
"train.csv",
transform=transforms.Compose([transforms.ConvertImageDtype(torch.float)]),
)
test_data = TrackManiaDataset(
"data",
"test.csv",
transform=transforms.Compose([transforms.ConvertImageDtype(torch.float)]),
)
return training_data, test_data
def load_dataloaders(training_data, test_data, batch_size=64):
train_dataloader = DataLoader(
training_data,
batch_size=batch_size,
shuffle=True,
num_workers=4,
pin_memory=True,
)
test_dataloader = DataLoader(
test_data, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True
)
return train_dataloader, test_dataloader
if __name__ == "__main__":
training_data, test_data = load_data()
view_data(training_data)
training_dataloader, _ = load_dataloaders(training_data, test_data)
view_dataloader(training_dataloader)