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

Tensorflow拼接张量单个元素时出现ZeroDivisionError报错问题

问题根因

这个错误是TensorFlow 2.7 CPU版本的tf.concat梯度实现缺陷导致的,两段代码的实际运算路径存在关键差异:

  • 直接相加的代码中,你操作的始终是形状为(1,)的变量t,t + t属于常规元素级运算,梯度传播逻辑简单无异常。
  • 触发错误的代码中,你先通过t[0]索引操作,把原本形状为(1,)的变量提取为0维标量张量,再将两个0维标量传入tf.concat并指定axis=0拼接。TensorFlow 2.7的tf.concat梯度算子在处理0维输入的梯度回传时,会错误地对拼接轴的维度长度做除法运算,而0维张量的对应轴长度为0,直接触发了除零错误。

修复方案

有两种方法可以解决这个问题:

  • 调整拼接输入的维度:在拼接前给0维标量增加维度,确保传入tf.concat的是1维张量,即可正常计算得到梯度值2.0,和直接相加的结果一致。修复后的示例代码如下:
import tensorflow as tf

with tf.GradientTape() as tape:
    t = tf.Variable([1.])
    # 给标量增加第0维后再拼接
    a = tf.expand_dims(t[0], axis=0)
    concat = tf.concat(values=[a, a], axis=0)
    concat_sum = tf.reduce_sum(concat)

grads = tape.gradient(concat_sum, t)
print(grads) # 输出 tf.Tensor([2.], shape=(1,), dtype=float32)
  • 升级TensorFlow版本:这个bug在TensorFlow 2.8及后续版本已经被官方修复,升级后原有代码无需修改即可正常运行。

内容的提问来源于stack exchange,提问作者jakob

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 06:45:03