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)
Logo

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

更多推荐