【发布时间】:2021-06-23 22:47:48
【问题描述】:
【问题讨论】:
-
你为什么要在内存中加载这么多数据?
-
用于训练深度学习模型
-
请使用数据生成器而不是一次加载所有数据
标签: python performance deep-learning google-colaboratory
【问题讨论】:
标签: python performance deep-learning google-colaboratory
一种高效且主要的方式是使用机器学习框架的数据加载器,例如 tensorflow、pytorch。根据您的代码,您当时正在加载所有图像,这需要很多时间。如果有很多图像,那么您可能会得到MemoryError。我强烈建议你在 PyTorch、Tensorflow 中使用 DataLoader。 DataLoaders 在训练过程中加载一批数据。
在 TensorFlow doc 中,您可以使用以下结构:
tf.keras.preprocessing.image_dataset_from_directory(
directory, labels='inferred', label_mode='int',
class_names=None, color_mode='rgb', batch_size=32, image_size=(256,
256), shuffle=True, seed=None, validation_split=None, subset=None,
interpolation='bilinear', follow_links=False
)
在 PyTorch doc 中,但这里首先需要指定数据集,然后将其交给数据加载器:
imagenet_data = torchvision.datasets.ImageNet('path/to/imagenet_root/')
data_loader = torch.utils.data.DataLoader(imagenet_data,
batch_size=4,
shuffle=True,
num_workers=args.nThreads)
上面提到的数据加载方法由于效率而被广泛使用。我希望使用它们可以帮助您完成任务。
【讨论】: