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

使用tf.GradientTape计算参与其他变量赋值的变量梯度

TensorFlow Eager模式下tf.GradientTape计算赋值关联变量梯度的方法

问题背景

在TensorFlow Eager模式下,需要计算损失相对于参与变量赋值运算的源变量的梯度,但直接调用gradient接口会返回None。
已查阅的历史相关资料中,一类方案未明确给出该场景的解决路径,另一类针对TensorFlow v1版本的方案仅适用于同一变量重复使用的场景,和当前需求不匹配;另有资料提到tf.assign类操作本身不支持梯度传播,仅给出张量层面的解决思路,无法直接落地到神经网络内部权重更新场景。

基础复现场景

import tensorflow as tf
a = tf.Variable(1.0, name='a')
b = tf.Variable(2.0, name='b')
c = tf.Variable(3.0, name='c')

with tf.GradientTape() as tape:
  c.assign(a + b)
  loss = tf.reduce_mean(c**2)

print(tape.gradient(loss, b)) # 输出None

# 手动watch变量的尝试也无效
with tf.GradientTape(watch_accessed_variables=False) as tape:
   tape.watch([b,c])
   c.assign(a + b)
   loss = tf.reduce_mean(c**2)

print(tape.gradient(loss, b)) # 仍然输出None

# 不使用assign、直接用张量运算可以正常得到梯度,但不符合变量赋值的使用需求
with tf.GradientTape() as tape:
   c = a + b
   loss = tf.reduce_mean(c**2)

print(tape.gradient(loss, b)) # 正常输出梯度值

嵌套变量扩展测试场景

针对列表嵌套变量的场景做了测试,代码如下:

import tensorflow as tf
a = [tf.Variable(1.0, name='a'), tf.Variable(4.0, name='aa')]
b = [tf.Variable(2.0, name='b'), tf.Variable(9.0, name='bb')]
c = [tf.Variable(3.0, name='c'), tf.Variable(0.0, name='cc')]
x = tf.Variable(0.01)

with tf.GradientTape(persistent=True) as tape:
    c_ = tf.nest.map_structure(lambda _a, _b: (1-x)*_a+ x*_b, a, b)
    tf.nest.map_structure(lambda x, y: x.assign(y), c, c_)
    loss = tf.norm(c) # 标量损失

# 该方式可以得到b对应的梯度
print(tape.gradient(loss,c,output_gradients=tape.gradient(c_,b)))
# [<tf.Tensor: shape=(), dtype=float32, numpy=0.0024197185>, <tf.Tensor: shape=(), dtype=float32, numpy=0.009702832>]
# 尝试获取x对应的梯度,结果为逐元素值,无法直接用于梯度下降
print(tape.gradient(loss,c,output_gradients=tape.gradient(c_,x)))
# [<tf.Tensor: shape=(), dtype=float32, numpy=1.4518311>, <tf.Tensor: shape=(), dtype=float32, numpy=5.8216996>]

# 不使用assign、直接对中间张量c_计算损失,可以正常得到x的梯度
with tf.GradientTape() as tape:
  c_ = tf.nest.map_structure(lambda _a, _b: (1-x)*_a+ x*_b, a, b)
  loss = tf.norm(c_) # 标量损失

print(tape.gradient(loss,x)) 
# tf.Tensor(5.0933886, shape=(), dtype=float32)

多维异构嵌套变量测试场景

进一步测试了多维度、不同形状嵌套变量的复杂场景,代码如下:

import tensorflow as tf

a = [tf.Variable([1.0, 2.0], name='a'), tf.Variable([5.0], name='aa'), tf.Variable(7.0, name='aaa')]
b = [tf.Variable([3.0, 4.0], name='b'), tf.Variable([6.0], name='bb'), tf.Variable(8.0, name='aaa')]
c = [tf.Variable([1.0, 1.0], name='c'), tf.Variable([1.0], name='cc'), tf.Variable(1.0, name='ccc')]
x = tf.Variable(0.5, name='x')

with tf.GradientTape(persistent=True) as tape:
    c_ = tf.nest.map_structure(lambda _a, _b: (1-x)*_a+ x*_b, a, b)

    tf.nest.map_structure(lambda x, y: x.assign(y), c, c_)

    loss = tf.norm(tf.nest.map_structure(lambda e: tf.norm(e), c))
    loss_without_assign = tf.norm(tf.nest.map_structure(lambda e: tf.norm(e), c_))

print(loss, loss_without_assign)
# tf.Tensor(9.974969, shape=(), dtype=float32) tf.Tensor(9.974969, shape=(), dtype=float32)

# 手动链式求导得到的结果和无assign场景的结果接近
#partial_grads = tf.nest.map_structure(lambda d, e: tf.nest.map_structure(lambda f, g: tape.gradient(loss, f, output_gradients=tape.gradient(g, x)), d, e), c, c_)
partial_grads = tf.nest.map_structure(lambda d, e: tape.gradient(loss, d, output_gradients=tape.gradient(e, x)), c, c_)

# 聚合逐元素梯度
print(tf.reduce_sum(tf.nest.map_structure(lambda z: tf.reduce_mean(z), partial_grads)))
print(tape.gradient(loss_without_assign, x))
# 结果接近:
# tf.Tensor(2.3057716, shape=(), dtype=float32)
# tf.Tensor(2.3057709, shape=(), dtype=float32)

根本原因

tf.Variable.assign()是状态写入操作,不属于TensorFlow自动微分追踪的可微张量运算范畴:自动微分仅记录张量之间的计算依赖关系,assign操作修改变量存储值的动作不会被记录到计算图中,因此无法通过assign建立源变量(比如示例中的a、b、x)和后续基于c计算的损失之间的依赖链,最终梯度返回None。

可行落地方案

针对需要保留变量c、同时要正确回传梯度到源变量的场景(比如神经网络权重原地更新的需求),采用链式求导+中间张量桥接的方式实现,步骤如下:

  • 先计算待赋值给目标变量的中间张量(比如示例中的c_),保留该张量上的全部计算依赖关系
  • 调用assign完成目标变量的原地更新,满足变量状态修改的业务逻辑
  • 分两段计算梯度:首先计算损失对被赋值变量c的梯度,再计算中间张量c_对源变量的梯度,通过gradient接口的output_gradients参数完成链式法则的梯度传递,最后对同源的梯度结果做聚合,即可得到正确的全局梯度。

该方案完全兼容嵌套结构、异构形状变量的场景,测试结果和无assign直接计算的梯度误差在浮点精度允许范围内,可以直接用于梯度下降优化流程。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 09:12:19