食物分类2
·
1.class semiDataset(Dataset),初始化是需要用到get_label()方法,x,y=get_label()
取出x的图片信息,y标签信息。if x中无数据,把flag设为false
else x中有数据,即半监督数据集中有数据,则flag=True,把x存入X中(self.X = np.array(x) ) 也存Y ,self.Y = torch.LongTensor(y)。
2.get_label():输入参数(无标签的数据集,模式,device,thres置信度)
把模型放到显卡上运行(加速),存储几个参数(置信度,标签,x,y)[]。
遍历无标签的数据(放到显卡上加速,计算预测,用sofemax计算概率,用 pred_max, pred_value分别记录每行的最大概率和对应的下标),
记录所有的pred_max, pred_value,然后遍如果pred_max>置信度,加入数据集
3.getitem(self, item): return self.transform(self.X[item]), self.Y[item]
class semiDataset(Dataset):
def __init__(self, no_label_loder, model, device, thres=0.99):
x, y = self.get_label(no_label_loder, model, device, thres)
if x == []:
self.flag = False #如果一张符合条件的图都没有 → flag=False
else:
self.flag = True
self.X = np.array(x) #存图片 # 转成 numpy 数组
self.Y = torch.LongTensor(y) #存标签
self.transform = train_transform #记录数据增强 train_transform
def get_label(self, no_label_loder, model, device, thres):
model = model.to(device)
pred_prob = [] # 存每个图片的置信度(模型有多确定)
labels = [] # 存模型预测的类别
x = [] # 最后留下的图片
y = [] # 最后留下的标签
soft = nn.Softmax() # Softmax回归,把输出转成 0~1 概率
with torch.no_grad():
for bat_x, _ in no_label_loder:
bat_x = bat_x.to(device)
pred = model(bat_x)
pred_soft = soft(pred)
pred_max, pred_value = pred_soft.max(1) #pred_max 记录的是概率,pred_value 记录的是下标
pred_prob.extend(pred_max.cpu().numpy().tolist()) #记录与预测值
labels.extend(pred_value.cpu().numpy().tolist()) #记录对应下标
for index, prob in enumerate(pred_prob):
if prob > thres: #如果预测值大于0.99就添加到x,y
x.append(no_label_loder.dataset[index][1]) #调用到原始的getitem
y.append(labels[index])
return x, y
def __getitem__(self, item):
return self.transform(self.X[item]), self.Y[item]
def __len__(self):
return len(self.X)更多推荐


所有评论(0)