You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使以下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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 06:18:53