代码之家  ›  专栏  ›  技术社区  ›  dragster

tensorflow的tf是怎样的。contrib。训练batch_sequences_with_states API工作?

  •  2
  • dragster  · 技术社区  · 8 年前

    我正在处理必须传递给RNN的长序列数据。要进行截断BPTT和批处理,似乎有两种选择:

    1. 通过组合创建批次 分别的 来自不同序列的片段。在一批中保留每个序列的最终状态,并将其传递给下一批。
    2. 将每个序列视为一个小批次,序列中的片段将成为该批次的成员。保留一段中最后一个时间步长的状态,并将其传递到下一段的第一个时间步长。

    我偶然发现 tf.contrib.training.batch_sequences_with_states

    我猜是第一种方式。这是因为,如果以第二种方式进行批处理,那么我们无法利用矢量化的好处,因为为了保持一个段的最后一个时间步长到下一个段的第一个时间步长之间的状态,RNN应该按顺序一次处理一个令牌。

    这两种批处理策略中的哪一种是在 tf。contrib。训练带状态的batch\u sequences\u ?

    1 回复  |  直到 8 年前
        1
  •  2
  •   Eugene Brevdo    8 年前

    tf.contrib.training.batch_sequences_with_states 实现前一个行为。每个小批量条目都是来自不同序列的一个段(每个序列可以由数量可变的段组成,具有唯一的键,该键被传递到 batch_sequences_with_states ). 与一起使用时 state_saving_rnn sess.run . 最后的片段为不同的序列腾出一个小批量插槽。