【问题标题】:Tensorflow: Can't understand ctc_beam_search_decoder() output sequenceTensorflow:无法理解 ctc_beam_search_decoder() 输出序列
【发布时间】:2017-08-03 11:24:55
【问题描述】:

我正在使用 Tensorflow 的 tf.nn.ctc_beam_search_decoder() 对 RNN 的输出进行解码,执行一些多对多映射(即每个网络单元的多个 softmax 输出)。

网络输出和 Beam 搜索解码器的简化版本是:

import numpy as np
import tensorflow as tf

batch_size = 4
sequence_max_len = 5
num_classes = 3

y_pred = tf.placeholder(tf.float32, shape=(batch_size, sequence_max_len, num_classes))
y_pred_transposed = tf.transpose(y_pred,
                                 perm=[1, 0, 2])  # TF expects dimensions [max_time, batch_size, num_classes]
logits = tf.log(y_pred_transposed)
sequence_lengths = tf.to_int32(tf.fill([batch_size], sequence_max_len))
decoded, log_probabilities = tf.nn.ctc_beam_search_decoder(logits,
                                                           sequence_length=sequence_lengths,
                                                           beam_width=3,
                                                           merge_repeated=False, top_paths=1)

decoded = decoded[0]
decoded_paths = tf.sparse_tensor_to_dense(decoded)  # Shape: [batch_size, max_sequence_len]

with tf.Session() as session:
    tf.global_variables_initializer().run()

    softmax_outputs = np.array([[[0.1, 0.1, 0.8], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1]],
                                [[0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]],
                                [[0.1, 0.7, 0.2], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]],
                                [[0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]]])

    decoded_paths = session.run(decoded_paths, feed_dict = {y_pred: softmax_outputs})
    print(decoded_paths)

这种情况下的输出是:

[[0]
 [1]
 [1]
 [1]]

我的理解是输出张量的维度应该是[batch_size, max_sequence_len],每一行都包含找到的路径中相关类的索引。

在这种情况下,我希望输出类似于:

[[2, 0, 0, 0, 0],
 [2, 2, 2, 2, 2],
 [1, 2, 2, 2, 2],
 [2, 2, 2, 2, 2]]

我对@9​​87654326@ 的工作原理有什么不明白的地方?

【问题讨论】:

    标签: python tensorflow beam-search


    【解决方案1】:

    tf.nn.ctc_beam_search_decoder documentation 所示,输出的形状不是[batch_size, max_sequence_len]。相反,它是

    [batch_size, max_decoded_length[j]]
    

    (在您的情况下使用j=0)。

    基于this paper 的第2 部分的开头(在github repository 中引用),max_decoded_length[0] 从上方以max_sequence_len 为界,但它们不一定相等。相关引文为:

    令 S 是一组从固定分布中抽取的训练样例 D_{XxZ}。输入空间 X = (R^m) 是 m 的所有序列的集合 维实值向量。目标空间 Z = L* 是 标签的(有限)字母 L 上的所有序列。一般来说,我们 将 L* 的元素称为标签序列或标签。每个例子 在 S 中由一对序列 (x, z) 组成。目标序列 z = (z1, z2, ..., zU) 最多与输入序列一样长 x = (x1, x2, ..., xT ),即 U

    其实max_decoded_length[0]依赖于具体的矩阵softmax_outputs。特别是,两个具有完全相同维度的矩阵可能会导致不同的max_decoded_length[0]

    例如,如果你替换行

    softmax_outputs = np.array([[[0.1, 0.1, 0.8], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1], [0.8, 0.1, 0.1]],
                                    [[0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]],
                                    [[0.1, 0.7, 0.2], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]],
                                    [[0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7], [0.1, 0.2, 0.7]]])
    

    与行

    np.random.seed(7)
    r=np.random.randint(0,100,size=(4,5,3))
    softmax_outputs=r/np.sum(r,2).reshape(4,5,1)
    

    你会得到输出

    [[1 0 1]
     [1 0 1]
     [1 0 0]
     [1 0 0]]
    

    (在上述示例中,softmax_outputs 由 logits 组成,并且与您提供的矩阵具有完全相同的维度)。

    另一方面,将种子更改为 np.random.seed(50) 会得到输出

    [[1 0]
     [1 0]
     [1 0]
     [0 1]]
    

    附言

    关于你问题的最后一部分:

    在这种情况下,我希望输出类似于:

    [[2, 0, 0, 0, 0],
     [2, 2, 2, 2, 2],
     [1, 2, 2, 2, 2],
     [2, 2, 2, 2, 2]]
    

    注意,基于documentationnum_classes实际上代表num_labels + 1。具体来说:

    输入张量的最内层维度大小,num_classes,表示 num_labels + 1 类,其中num_labels 是真实标签的数量, 最大值 (num_classes - 1) 保留给空白 标签。

    例如,对于包含 3 个标签 [a, b, c] 的词汇表, num_classes = 4 并且标签索引是 {a: 0, b: 1, c: 2, 空白: 3}。

    因此,在您的情况下,真正的标签是 0 和 1,而 2 是为空白标签保留的。空白标签代表观察无标签的情况(第3.1节here):

    CTC 网络有一个 softmax 输出层(Bridle,1990),还有一个 单位比 L 中有标签。第一个 |L| 的激活 单位被解释为观察到的概率 在特定时间对应的标签。 额外的激活 unit 是观察到“空白”或没有标签的概率。 这些输出定义了所有可能方式的概率 将所有可能的标签序列与输入序列对齐。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2016-12-07
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2017-02-01
      • 2013-12-07
      • 1970-01-01
      • 2021-03-08
      相关资源
      最近更新 更多