Pytorch中的dataset类,应用于Siamese network
转自https://blog.csdn.net/leviopku/article/details/99958182
作者写的非常清楚,我是pytorch小白,对于定义dataset类很迷茫,官方文档对新手很不友好,看了作者的博客豁然开朗。
dataset类,创建适应模型的数据接口
官方给出标准,在创建dataset类时,必须有__getitem__和__len__。__getitem__就是获取样本对,模型直接通过这一函数获得一对样本对{x:y},__len__是指数据集长度。
例如,自己新建一个dataset:
1 class MyDataSet(Dataset): 2 def __init__(self): 3 self.sample_list = ... 4 5 def __getitem__(self, index): 6 x= ... 7 y= ... 8 return x, y 9 10 def __len__(self): 11 return len(self.sample_list)
但有时候,我们使用自己的数据时,标签和数据不在一个数据中,可以通过一个txt文件实现映射
例如,我的模型是Siamese network,输入需要(input1,input2,label)形式,创建自己的dataset如下:
from torch.utils.data import Dataset class SiameseNetworkDataset(): def __init__(self, data_dir, transform=None): self.data_dir = data_dir dataset_path = '/CERT_Dataset/Data_process/Payrall' self.transform = transform self.sample_list = list() # self.dataset_type = dataset_type f = open(data_dir + '/datalist.txt') lines = f.readlines() for line in lines: self.sample_list.append(line.strip()) f.close() def __getitem__(self,index): data0_index = np.random.randint(0, len(self.sample_list)) data0_path = self.sample_list[data0_index].split(',')[0] data0 = pd.read_csv(data0_path) label0 = int(self.sample_list[data0_index].split(',')[-1]) data1_path = self.sample_list[index].split(',')[0] data1 = pd.read_csv(data1_path) label1 = int(self.sample_list[index].split(',')[-1]) target = 1.0 if label0 == label1 else 0.0 if self.transform is not None: data0 = self.transform(data0) data1 = self.transform(data1) return data0, data1, target def __len__(self): return len(self.sample_list)
测试dataset类是否成功:
if __name__ == '__main__': ds = SiameseNetworkDataset() print(ds.__len__()) data0,data1,target = ds.__getitem__(1) print(type(data0)) print(data0)

浙公网安备 33010602011771号