pytorch-训练自定义数据集实战

1. 步骤

  • 加载数据
  • 创建模型
  • 训练和测试
  • 迁移学习

2. 加载数据

这里以宝可梦动画图片为数据集
在这里插入图片描述
下载地址:
链接:https://pan.baidu.com/s/1TbXKNIBitXk_o-oVAAiX-A?pwd=py3r
提取码:py3r

数据集各分类情况和切分比例见下图:
在这里插入图片描述

2.1 继承Dataset

继承torch.utils.data.Dataset,实现__len__和__getitem__函数
__len__是获取所有数据集的数量
__getitem__获取数据集中指定index的image tensor和对应的分类label
实现这两个函数的思路:

  • 将数据集所有文件名加载到list中,通过len([images]),即可实现__len__
  • 生成数据集中所有的image path和label,读取并预处理image,即可实现__getitem__

2.1.1 生成name2label

数据集文件结构是pokemon\bulbasaur\00000000.png,pokemon下的每个文件夹代表一个分类,因此就可以实现下面的代码生成一个name2label

 self.name2label = {
   
   } # "sq...":0
 for name in sorted(os.listdir(os.path.join(root))):
     if not os.path.isdir(os.path.join(root, name)):
         continue

     self.name2label[name] = len(self.name2label.keys())

2.1.2 生成image path, label的文件

获取pokemon目前下所有数据文件的路径放到images中,遍历images,通过每条数据文件路径中的分类文件夹名称从name2label获取到对应的label,然后写入到文件中。
代码如下:

       if not os.path.exists(os.path.join(self.root, filename)):
            images = []
            for name in self.name2label.keys():
                # 'pokemon\\mewtwo\\00001.png
                images += glob.glob(os.path.join(self.root, name, '*.png'))
                images += glob.glob(os.path.join(self.root, name, '*.jpg'))
                images += glob.glob(os.path.join(self.root, name, '*.jpeg'))

            # 1167, 'pokemon\\bulbasaur\\00000000.png'
            print(len(images), images)

            random.shuffle(images)
            with open(os.path.join(self.root, filename), mode='w', newline='') as f:
                writer = csv.writer(f)
                for img in images: # 'pokemon\\bulbasaur\\00000000.png'
                    name = img.split(os.sep)[-2]
                    label = self.name2label[name]
                    # 'pokemon\\bulbasaur\\00000000.png', 0
                    writer.writerow([img, label])
                print('writen into csv file:', filename)

2.1.3 len

    def __len__(self):

        return len(self.images)

2.1.3 getitem

预处理包括resize、randomRotation、ToTensor、Normalize等

def __getitem__(self, idx):
    # idx~[0~len(images)]
      # self.images, self.labels
      # img: 'pokemon\\bulbasaur\\00000000.png'
      # label: 0
      img, label = self.images[idx], self.labels[idx]

      tf = transforms.Compose([
          lambda x:Image.open(x).convert('RGB'), # string path= > image data
          transforms.Resize((int(self.resize*1.25), int(self.resize*1.25))),
          transforms.RandomRotation(15),
          transforms.CenterCrop(self.resize),
          transforms.ToTensor(),
          transforms.Normalize(mean=[0.485, 0.456, 0.406],
                               std=[0.229, 0.224, 0.225])
      ])

      img = tf(img)
      label = torch.tensor(label)


      return img, label

2.1.4 数据切分为train、val、test

数据切分比例6:2:2

class Pokemon(Dataset):

    def __init__(self, root, resize, mode):
        super(Pokemon, self).__init__()

        self.root = root
        self.resize = resize

        self.name2label = {
   
   } # "sq...":0
        for name in sorted(os.listdir(os.path.join(root))):
            if not os.path.isdir(os.path.join(root, name)):
                continue

            self
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值