关于tf.nn.raw_rnn中emit_structure为None引发tf.where报错的技术疑问
关于
tf.nn.raw_rnn中emit_structure为None的报错问题 你提到的这个问题确实是使用tf.nn.raw_rnn时很容易踩的一个坑,官方文档这里的描述确实不够明确,容易让人产生困惑!
问题根源
按照文档逻辑,首次运行loop_fn时返回的emit_structure,是用来给后续已完成的小批量条目生成占位零张量的——也就是通过emit = tf.where(finished, tf.zeros_like(emit_structure), emit)这条语句,把已完成条目的输出替换成和emit_structure同形状同类型的零。但如果emit_structure设为None,tf.zeros_like(emit_structure)必然会抛出错误,因为tf.zeros_like的输入必须是一个有效的张量,None完全不满足要求。
解决思路
这里给你两个可行的解决方向:
- 返回合法的零张量作为初始
emit_structure:如果你的初始步骤没有实际输出,不要返回None,而是返回一个和后续emit形状、数据类型完全一致的零张量。比如后续emit是形状为[batch_size, 64]的float32张量,那首次loop_fn返回的emit_structure就可以写成:
这样后续的emit_structure = tf.zeros((batch_size, 64), dtype=tf.float32)tf.where就能正常生成对应的占位零张量,保证输出形状统一。 - 调整
loop_fn逻辑避免依赖emit_structure:如果你的场景中不需要给已完成条目填充零(这种情况比较少见,因为批量处理通常需要输出形状一致),可以重新设计loop_fn的输出逻辑,去掉依赖emit_structure的tf.where语句,但要注意保证所有批量条目的输出长度对齐。
内容的提问来源于stack exchange,提问作者figs_and_nuts
相关产品推荐
相关产品推荐

