自定义数据集
| 5 | from torch.utils.data import Dataset |
| 6 | |
| 7 | class MyDataSet(Dataset): |
| 8 | """自定义数据集""" |
| 9 | def __init__(self, images_path, images_class, 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 __gettiem__(self, item): |
| 18 | img = Image.open(self.images_path[item]) |
| 19 | if img.mode != 'RGB': |
| 20 | raise ValueError("image: {} isn't RGB mode.".format(self.images_path[item])) |
| 21 | label = self.images_class[item] |
| 22 | |
| 23 | if self.transform is not None: |
| 24 | img = self.transform(img) |
| 25 | return img, label |
| 26 | |
| 27 | @staticmethod |
| 28 | def collate_fn(batch): |
| 29 | images, label = tuple(zip(*batch)) |
| 30 | images = torch.stack(images, dim=0) |
| 31 | labels = torch.as_tensor(labels) |
| 32 | return images, labels |
| 33 | |
| 34 |