你说得对,docs 在这方面不是很明白...
您需要的实际上是一个 4 步过程:
- 定义数据增强
- 适合增强
- 使用
flow_from_directory() 设置您的生成器
- 使用
fit_generator() 训练您的模型
这是假设图像分类案例的必要代码:
# define data augmentation configuration
train_datagen = ImageDataGenerator(featurewise_center=True,
featurewise_std_normalization=True,
zca_whitening=True)
# fit the data augmentation
train_datagen.fit(x_train)
# setup generator
train_generator = train_datagen.flow_from_directory(
train_data_dir,
target_size=(img_height, img_width),
batch_size=batch_size,
class_mode='categorical')
# train model
model.fit_generator(
train_generator,
steps_per_epoch=nb_train_samples,
epochs=epochs,
validation_data=validation_generator, # optional - if used needs to be defined
validation_steps=nb_validation_samples)
显然,有几个参数需要定义(train_data_dir、nb_train_samples 等),但希望你能明白。
如果您还需要使用validation_generator,如我的示例所示,则应以与train_generator 相同的方式定义。
更新(评论后)
第 2 步需要一些讨论;这里,x_train 是实际数据,理想情况下,它应该适合主内存。还有(documentation),这一步是
仅当 featurewise_center 或 featurewise_std_normalization 或 zca_whitening 时才需要。
但是,在许多实际案例中,要求所有训练数据都放入内存显然是不现实的。在这种情况下如何对数据进行中心化/标准化/白色化本身就是一个(巨大的)子领域,并且可以说是 Spark 等大数据处理框架存在的主要原因。
那么,在实践中该怎么做呢?那么,在这种情况下,下一个合乎逻辑的操作是采样您的数据;事实上,这正是社区所建议的——这里是 Keras 的创造者 Francois Chollet,Working with large datasets like Imagenet:
datagen.fit(X_sample) # let's say X_sample is a small-ish but statistically representative sample of your data
还有来自ongoing open discussion 的关于扩展ImageDataGenerator 的另一句话(强调):
fit 是 feature-wise 标准化和 ZCA 所必需的,它只需要一个数组作为参数,不适合目录。 现在,我们需要手动读取图像的一个子集以适应目录。一个想法是我们可以更改fit() 以接受生成器本身(flow_from_directory),当然,在拟合期间应该禁用标准化。