6 python database hdf5 deep-learning pytorch
我正在尝试训练深度学习模型,而不将整个数据集加载到内存中。我的主要问题是,这样做的最佳方法是什么?
HDF5 似乎是人们实现此目的的常用方法,也是我首先尝试的。然而,当使用 pytorch 的 dataloader 类时,运行速度非常慢。我创建了自己的迭代器,它运行得更快,但是数据并不是每批都是随机的。我试图了解为什么 pytorch 数据加载器运行缓慢以及我是否可以对此做些什么。
下面是我的代码
首先,我定义了一个数据集类,它接受 HDF5 数据集的文件路径。我对此代码的理解是,每当调用getitem时,它都会从磁盘读取。
class My_H5Dataset(torch.utils.data.Dataset):
def __init__(self, file_path):
super(My_H5Dataset, self).__init__()
h5_file = h5py.File(file_path , 'r')
self.features = h5_file['features']
self.labels = h5_file['labels']
self.index = h5_file['index']
self.labels_values = h5_file['labels_values']
self.index_values = h5_file['index_values']
def __getitem__(self, index):
return (torch.from_numpy(self.features[index,:]).float(),
torch.from_numpy(self.labels[index,:]).float(),
torch.from_numpy(self.index[index,:]).float(),
torch.from_numpy(np.array(self.labels_values[index])),
torch.from_numpy(np.array(self.index_values[index])))
def __len__(self):
return self.features.shape[0]
Run Code Online (Sandbox Code Playgroud)
然后我简单地将其传递到 pytorch 数据加载器中,如下所示
train_dataset = My_H5Dataset(hdf5_data_folder_train)
train_ms = MySampler(train_dataset)
trainloader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size,
sampler=train_ms,num_workers=2)
Run Code Online (Sandbox Code Playgroud)
我的另一种方法是手动定义迭代器。而且这确实运行得更快。然而,我没有注意到 HDD 与 SSD 中存储的数据之间的速度差异,这让我担心我遗漏的某个地方存在瓶颈。
def hdf5_loader_generator(dataset, batch_size, as_tensor=True, n_samples = 10000):
"""Given an h5 path to a file that holds the arrays, returns a generator
that can get certain data at a time."""
stop = n_samples
curr_index = start = 0
while 1:
stop_index = min([curr_index + batch_size, stop])
ft, sp, index, sp_values, index_values = dataset[curr_index:stop_index]
curr_index += batch_size
if curr_index >= stop:
curr_index = start
continue
yield ft, sp, index, sp_values, index_values
def hdf5_data_iterator(dataset, batch_size, as_tensor=True):
return iter(hdf5_loader_generator(dataset, batch_size, as_tensor))
Run Code Online (Sandbox Code Playgroud)
我对 pytorch 迭代器的计时如下:
start = time.time()
for i, data in enumerate(trainloader, 0):
ft, sp, index, sp_values, index_values = data
ft, sp, index = ft.to(device), sp.to(device), index.to(device)
if i > 500:
break
end = time.time()
print(end-start)
Run Code Online (Sandbox Code Playgroud)
我对迭代器的计时如下:
train_dataset = My_H5Dataset(hdf5_data_folder_train)
samples = hdf5_data_iterator(train_dataset, batch_sizes)
start = time.time()
for i, data in enumerate(samples):
ft, sp, index, sp_values, index_values = data
ft, sp, index = ft.to(device), sp.to(device), index.to(device)
if i > 500:
break
end = time.time()
print(end-start)
Run Code Online (Sandbox Code Playgroud)
有比我目前正在尝试的更好的方法吗?
感谢您的帮助!
| 归档时间: |
|
| 查看次数: |
1236 次 |
| 最近记录: |