【问题标题】:How to properly shuffle a dataset in Tensorflow after every epoch如何在每个 epoch 后正确打乱 Tensorflow 中的数据集
【发布时间】:2021-07-06 16:49:14
【问题描述】:

我目前正在使用 Tensorflow 和 Keras 开发一个神经网络,我有一个写在 TFRecord 上的数据集,我必须从中读取数据,问题是神经网络是在体积上训练的,我没有足够的内存将所有内容存储在 ram 中,我正在读取这样的数据,代码取自这两个地方:

https://colab.research.google.com/github/GoogleCloudPlatform/training-data-analyst/blob/master/courses/fast-and-lean-data-science/07_Keras_Flowers_TPU_solution.ipynb

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


    【解决方案1】:

    你可以看看Custom DataGenerators

    在他们的on_epoch_end 方法中,您可以随机播放数据,这些数据将在每个时期从您的内存中加载。

    【讨论】:

    • 感谢您的参考,我从来没有接触过生成器,我一直依赖 API,我会试一试,但我仍然不明白是什么问题
    【解决方案2】:

    您正在使用 dataset.shuffle(),然后执行 .cache()。由于您每次都在更改数据顺序,因此 tensorflow 会将每个经过洗牌的数据集缓存在内存中。这会导致同一数据集的多个混洗副本和 RAM 被填满。

    【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2017-10-22
    • 1970-01-01
    • 2019-01-16
    • 2020-01-19
    • 2018-05-08
    • 1970-01-01
    • 1970-01-01
    • 2018-03-16
    相关资源
    最近更新 更多