食物分类1
1.初始化数据,传入路径和模式
2.判断模式,if半监督  则只提取图片信息 else 提取图片信息和标签信息,并转化为长整型

3.读数据函数:if半监督  读取文件夹里所有图片,并建立一个可以存储图片信息的张量(Xi),遍历所读取的图片信息,把path,img组合起来,修改图片大小为固定值,把图片依次存入Xi中。

else (训练模式)先循环11次,生成文件名并列出所有文件名,建立一个可以存储图片信息的张量(Xi)和标签信息的(Yi),把path,img组合起来,修改图片大小为固定值,把图片依次存入Xi中。如果是第一张照片则传入Xi,Yi,如果不是第一张照片,把当前Xi,Yi与之前的X,Y结合在一起

4.getitem(self, item):if 半监督 返回transform(self.X[item]),X[item]
else 训练模式 self.transform(self.X[item]), self.Y[item]


class food_Dataset(Dataset):
    def __init__(self, path, mode="train"):       #如果你不传 mode,它就自动用 "train"!
        self.mode = mode
        if mode == "semi":                       #无标签的半监督模式   则只需要读取图片信息  ,无标签信息
            self.X = self.read_file(path)
        else:
            self.X, self.Y = self.read_file(path)
            self.Y = torch.LongTensor(self.Y)  #标签转为长整形\

        if mode == "train":
            self.transform = train_transform
        else:
            self.transform = val_transform

    def read_file(self, path):
        if self.mode == "semi":                     #semi 模式:只读图
            file_list = os.listdir(path)            # 读取文件夹里所有图片
            xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)    #(数量,高,宽,3) → 专门用来存一摞图片的数组!
            # 列出文件夹下所有文件名字
            for j, img_name in enumerate(file_list):
                img_path = os.path.join(path, img_name)       #"folder/00/a.jpg"  路径+jpg
                img = Image.open(img_path)
                img = img.resize((HW, HW))                  #把图片大小改为224*224
                xi[j, ...] = img                            #按xi[j]的顺序存储img
            print("读到了%d个数据" % len(xi))
            return xi
        else:
            for i in tqdm(range(11)):                           #tqdm显示循环进度条
                file_dir = path + "/%02d" % i                  # 进入 00 文件夹、01 文件夹...
                file_list = os.listdir(file_dir)              #列出所有文件的名字,

                xi = np.zeros((len(file_list), HW, HW, 3), dtype=np.uint8)
                yi = np.zeros(len(file_list), dtype=np.uint8)

                # 列出文件夹下所有文件名字
                for j, img_name in enumerate(file_list):
                    img_path = os.path.join(file_dir, img_name)
                    img = Image.open(img_path)
                    img = img.resize((HW, HW))              #固定长宽为224*224
                    xi[j, ...] = img                        # 存图片
                    yi[j] = i                           # 标签就是文件夹号 0、1、2...10

                if i == 0:                               #如果是第一张图片
                    X = xi
                    Y = yi
                else:
                    X = np.concatenate((X, xi), axis=0)                       #竖着合并
                    Y = np.concatenate((Y, yi), axis=0)
            print("读到了%d个数据" % len(Y))
            return X, Y

    def __getitem__(self, item):
        if self.mode == "semi":      #无标签数据 → 只给你两张一样的图,不给标签!
            return self.transform(self.X[item]), self.X[item]
        else:     #有标签数据 → 给你 图片 + 答案!
          return self.transform(self.X[item]), self.Y[item]

    def __len__(self):
        return len(self.X)
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐