自定义数据集
| 4 | |
| 5 | |
| 6 | class 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 |