如何使以下TensorFlow scan矩阵示例正常运行?用于测试函数f代码
解决TensorFlow Scan函数运行问题
首先,咱们来拆解下你的代码里的核心问题:
- Initializer维度不匹配:你给
tf.scan传入的initializer=0是个标量,但输入a里的每个元素都是长度为5的向量,而且f函数返回的是经过softmax处理后的同维度向量。tf.scan要求初始值的维度必须和函数返回值、输入元素的维度完全一致,否则会直接抛出维度不匹配的错误。 - 另外,虽然你的
f函数目前没用到prev_y参数,但tf.scan的回调函数必须固定接收两个参数(前一次的输出结果、当前输入元素),这个参数可以保留但不用它,完全没问题。
接下来是修改后的可运行代码:
import tensorflow as tf def f(prev_y, curr_y): # prev_y参数保留即可,不用也没关系 fval = tf.nn.softmax(curr_y) return fval a = tf.constant([[.1, .25, .3, .2, .15], [.07, .35, .27, .17, .14]]) # 将initializer改为和curr_y同维度的零向量,长度为5 c = tf.scan(f, a, initializer=tf.zeros([5])) with tf.Session() as sess: print(sess.run(c))
运行这段代码后,你会得到每个输入向量经过softmax处理后的结果——因为这里的scan只是逐个处理每个元素(没用到prev_y),效果和直接对a做softmax类似,但如果后续你需要用到前一次的输出做累积计算(比如累加概率、递推更新),只需要在f函数里补充逻辑就行,比如:
def f(prev_y, curr_y): # 累积前一次结果与当前softmax值 fval = prev_y + tf.nn.softmax(curr_y) return fval
内容的提问来源于stack exchange,提问作者John
相关产品推荐
相关产品推荐

