【问题标题】:Tensorflow - LSTM state reuse within batchTensorflow - 批次内的 LSTM 状态重用
【发布时间】:2017-02-09 10:06:30
【问题描述】:

我正在研究使用 LSTM 跟踪参数的 Tensorflow NN(时间序列数据回归问题)。一批训练数据包含 连续 个观察值的 batch_size。我想使用 LSTM 状态作为下一个样本的输入。因此,如果我有一批数据观察,我想将第一个观察的状态作为输入提供给第二个观察,依此类推。下面我将 lstm 状态定义为 size = batch_size 的张量。我想批量重用状态within:

state = tf.Variable(cell.zero_states(batch_size, tf.float32), trainable=False)
cell = tf.nn.rnn_cell.BasicLSTMCell(100)
output, curr_state = tf.nn.rnn(cell, data, initial_state=state) 

在 API 中有一个 tf.nn.state_saving_rnn 但文档有点模糊。 我的问题:如何在训练批次中重复使用 curr_state。

【问题讨论】:

  • 为了澄清,您想将第一个批处理元素的结果中的状态线程化为下一个批处理元素的开始状态,等等?在这种情况下,批次维度不正是时间维度吗?
  • @Allen Lavoie,是的,没错。批次中的每个数据观察都是一个(多维)时间序列窗口。该批次包含按顺序排列的重叠窗口。 batch维度是时间维度,有重叠和跨步。
  • 在这种情况下,您的批处理维度实际上是 1。除非您有多个序列可以一起批处理,否则这将相对较慢。正在努力支持近似值,允许对单个较长时间序列进行批处理,但尚未公开发布任何内容。
  • 感谢您的解释!如果您能详细解释一下“为单个较长时间序列进行批处理”的工作原理并写下答案,我会标记它。

标签: tensorflow lstm recurrent-neural-network


【解决方案1】:

你基本上就在那里,只需要将state更新为curr_state:

state_update = tf.assign(state, curr_state)

然后,确保您在 state_update 本身上调用 run 或将 state_update 作为依赖项的操作,否则分配实际上不会发生。例如:

with tf.control_dependencies([state_update]):
    model_output = ...

正如 cmets 中所建议的,RNN 的典型情况是您有一个批次,其中第一个维度 (0) 是序列数,第二个维度 (1) 是每个序列的最大长度(如果您通过time_major=True 当你构建 RNN 时,这两个被交换)。理想情况下,为了获得良好的性能,您可以将多个序列堆叠成一个批次,然后按时间拆分该批次。但这真的是一个不同的话题。

【讨论】:

    猜你喜欢
    • 2017-05-28
    • 2016-11-09
    • 2020-02-03
    • 2018-10-07
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2017-09-28
    相关资源
    最近更新 更多