图像分类任务
·
食物分类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)更多推荐


所有评论(0)