由于您计划使用可迭代数据集,因此您不需要随机访问(IterableDataset 不支持随机采样器)。在这种情况下,为什么不将所有内容都写入二进制文件并对其进行迭代呢?我发现在实践中这通常比替代解决方案快得多。这应该比保存为文本文件快得多,因为您避免了将文本转换为数字的开销。
示例实现可能如下所示。首先我们可以如下构建一个二进制文件(包含随机数据作为占位符)
import numpy as np
from tqdm import tqdm
filename = 'data.bin'
num_samples = 3600000
rows, cols = 30, 32
dtype = np.float32
# format: <num_samples> <rows> <cols> <sample0> <sample1>...
with open(filename, 'wb') as fout:
# write a header that contains the total number of samples and the rows and columns per sample
fout.write(np.array((num_samples, rows, cols), dtype=np.int32).tobytes())
for i in tqdm(range(num_samples)):
# random placeholder
sample = np.random.randn(rows, cols).astype(dtype)
# write data to file
fout.write(sample.tobytes())
那么我们可以定义一个IterableDataset如下
import numpy as np
from torch.utils.data import IterableDataset, DataLoader
from tqdm import tqdm
def binary_reader(filename, start=None, end=None, dtype=np.float32):
itemsize = np.dtype(dtype).itemsize
with open(filename, 'rb') as fin:
num_samples, rows, cols = np.frombuffer(fin.read(3 * np.dtype(np.int32).itemsize), dtype=np.int32)
start = start if start is not None else 0
end = end if end is not None else num_samples
blocksize = itemsize * rows * cols
start_offset = start * blocksize
fin.seek(start_offset, 1)
for _ in range(start, end):
yield np.frombuffer(fin.read(blocksize), dtype=dtype).reshape(rows, cols).copy()
class BinaryIterableDataset(IterableDataset):
def __init__(self, filename, start=None, end=None, dtype=np.float32):
super().__init__()
self.filename = filename
self.start = start
self.end = end
self.dtype = dtype
def __iter__(self):
return binary_reader(self.filename, self.start, self.end, self.dtype)
通过在我的系统(使用 SSD 存储)上对该数据集的快速测试,我发现我能够在大约 10 秒内迭代所有 360 万个样本
dataset = BinaryIterableDataset('data.bin')
for sample in tqdm(dataset):
pass
3600000it [00:09, 374026.17it/s]
使用DataLoader 和batch_size=256 需要大约20 秒来迭代整个数据集(转换为张量和创建批次有一些开销)。对于这个数据集,我发现使用并行加载时将数据传入和传出共享内存的开销实际上比仅使用 0 个 worker 慢很多。因此我推荐使用num_workers=0。与任何可迭代数据集一样,您需要添加额外的逻辑来支持 num_workers > 1,尽管我不确定在这种情况下是否值得。
loader = DataLoader(dataset, batch_size=256, num_workers=0)
for batch in tqdm(loader):
# batch is a tensor of shape (256, 30, 32)
pass
14063it [00:19, 710.49it/s]
请注意,data.bin 文件不能跨使用不同字节顺序的系统移植。尽管可以进行修改以支持这一点。