【发布时间】:2021-07-06 16:49:14
【问题描述】:
我目前正在使用 Tensorflow 和 Keras 开发一个神经网络,我有一个写在 TFRecord 上的数据集,我必须从中读取数据,问题是神经网络是在体积上训练的,我没有足够的内存将所有内容存储在 ram 中,我正在读取这样的数据,代码取自这两个地方:
https://keras.io/examples/keras_recipes/tfrecord/
def load_dataset(filenames):
option_no_order = tf.data.Options()
option_no_order.experimental_deterministic = False
dataset = tf.data.TFRecordDataset(filenames, num_parallel_reads=AUTO)
dataset = dataset.with_options(option_no_order)
dataset = dataset.map(decode_record, num_parallel_calls=AUTO)
return dataset
def get_batched_dataset(filenames, train=False):
dataset = load_dataset(filenames)
if train:
dataset = dataset.repeat() # Best practices for Keras: Training dataset: repeat then batch Evaluation dataset: do not repeat
dataset = dataset.cache() # This dataset fits in RAM
dataset = dataset.batch(BATCH_SIZE)
dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)
return dataset
这段代码有效,但我想在每个 epoch 之后添加训练数据集的改组,我写了这个:
def get_batched_dataset(filenames, train=False):
dataset = load_dataset(filenames)
if train:
dataset = dataset.shuffle(200, reshuffle_each_iteration=True) ###############
dataset = dataset.repeat() # Best practices for Keras: Training dataset: repeat then batch Evaluation dataset: do not repeat
dataset = dataset.cache() # This dataset fits in RAM
dataset = dataset.batch(BATCH_SIZE)
dataset = dataset.prefetch(AUTO) # prefetch next batch while training (autotune prefetch buffer size)
return dataset
这段代码可以正常工作,但我观察到每个 epoch 之后 RAM 使用量都有所增加,在 4 个 epoch 之后,整个会话因“内存不足”而崩溃。
我得到数据集并像这样训练网络:
train_dataset = get_directories()
training_dataset = get_batched_dataset('train.tfrecords', train=True)
validation_dataset = get_batched_dataset('valid.tfrecords', train=False)
model.fit(training_dataset, steps_per_epoch=len(train_dataset), epochs=80, validation_data=validation_dataset, callbacks=my_callbacks)
shuffle 函数获取 200 个卷并将它们放入缓冲区并随机馈送到网络,我不明白为什么会话会这样崩溃
【问题讨论】:
标签: python tensorflow keras google-colaboratory