【问题标题】:What does train_data.cache().shuffle(BUFFER_SIZE).batch(BATCH_SIZE).repeat() do?train_data.cache().shuffle(BUFFER_SIZE).batch(BATCH_SIZE).repeat() 有什么作用?
【发布时间】:2020-09-04 10:02:15
【问题描述】:

我正在关注 Tensorflow 的时间序列/LSTM 教程,并且很难理解这条线的作用,因为它没有得到真正的解释:

train_data.cache().shuffle(BUFFER_SIZE).batch(BATCH_SIZE).repeat()

我试图查看不同模块的作用,但我无法理解完整的命令及其对数据集的影响。 这是整个教程: Click

【问题讨论】:

    标签: python tensorflow time-series lstm


    【解决方案1】:

    这是一个基于tensorflow.data API 的输入管道定义。 分解:

    (train_data # some tf.data.Dataset, likely in the form of tuples (x, y)
    .cache() # caches the dataset in memory (avoids having to reapply preprocessing transformations to the input)
    .shuffle(BUFFER_SIZE) # shuffle the samples to have always a random order of samples fed to the network
    .batch(BATCH_SIZE) # batch samples in chunks of size BATCH_SIZE (except the last one, that may be smaller)
    .repeat()) # repeat forever, meaning the dataset will keep producing batches and never terminate running out of data.
    

    注意事项:

    • 因为重复是在 shuffle 之后,所以批次总是不同的,即使跨 epoch 也是如此。
    • 由于cache(),数据集的第二次迭代将从内存中的缓存加载数据,而不是管道的先前步骤。如果数据预处理很复杂,这可以为您节省一些时间(但是,对于大型数据集,这可能会占用您的内存)
    • BUFFER_SIZE 是随机播放缓冲区中的项目数。该函数填充缓冲区,然后从中随机采样。适当的洗牌需要足够大的缓冲区,但这是与内存消耗的平衡。重新洗牌会在每个 epoch 自动发生。

    注意:这是一个管道定义,所以你要重新指定管道中的操作,而不是实际运行它们!这些操作实际上是在您调用 next(iter(dataset)) 时发生的,而不是之前。

    【讨论】:

    • 谢谢,但是我有两个问题:缓冲区大小到底有什么作用?我们不需要在 epoch 之间再次调用这条线,以便在它们之间重新洗牌吗?
    猜你喜欢
    • 2021-01-29
    • 1970-01-01
    • 1970-01-01
    • 2019-11-18
    • 2018-09-29
    • 1970-01-01
    • 2014-12-19
    • 1970-01-01
    • 2016-06-19
    相关资源
    最近更新 更多