MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / MyDataSet

Class MyDataSet

CV/Pytorch_classification/ShuffleNet/my_dataset.py:6–38  ·  view source on GitHub ↗

自定义数据集

Source from the content-addressed store, hash-verified

4
5
6class MyDataSet(Dataset):
7 """自定义数据集"""
8
9 def __init__(self, images_path: list, images_class: list, transform=None):
10 self.images_path = images_path
11 self.images_class = images_class
12 self.transform = transform
13
14 def __len__(self):
15 return len(self.images_path)
16
17 def __getitem__(self, item):
18 img = Image.open(self.images_path[item])
19 # RGB为彩色图片,L为灰度图片
20 if img.mode != 'RGB':
21 raise ValueError("image: {} isn't RGB mode.".format(self.images_path[item]))
22 label = self.images_class[item]
23
24 if self.transform is not None:
25 img = self.transform(img)
26
27 return img, label
28
29 @staticmethod
30 # 用法:将batch_size的数据进行整理,自定义数据的堆叠过程和batch数据的输出形式
31 def collate_fn(batch):
32 # 官方实现的default_collate可以参考
33 # https://github.com/pytorch/pytorch/blob/67b7e751e6b5931a9f45274653f4f653a4e6cdf6/torch/utils/data/_utils/collate.py
34 images, labels = tuple(zip(*batch))
35
36 images = torch.stack(images, dim=0)
37 labels = torch.as_tensor(labels)
38 return images, labels

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected