TensorFlow中何时需用watch()追踪梯度?两种场景差异解析
TensorFlow中tf.GradientTape.watch()的使用差异
核心区别在于被追踪的张量类型:
- tf.GradientTape默认会自动追踪**可训练变量(tf.Variable)**的梯度计算
- 普通张量(tf.Tensor)不会被自动追踪,必须手动调用
t.watch(x)才能让梯度带记录它的运算过程
第一个例子(需要watch())
这里的x是用tf.ones()创建的普通tf.Tensor,不属于可训练变量范畴。如果不在GradientTape上下文里调用t.watch(x),梯度带不会记录x参与的运算,后续计算dz_dx时会得到None。
x = tf.ones((2,2)) with tf.GradientTape() as t: # 手动让梯度带追踪普通张量x t.watch(x) y = tf.reduce_sum(x) z = tf.square(y) # 计算z对x的梯度 dz_dx = t.gradient(z, x)
第二个例子(不需要watch())
这里求导的目标是model.trainable_variables,这些都是tf.Variable类型的可训练变量(比如模型的权重、偏置参数)。GradientTape默认会自动追踪所有可训练变量的运算过程,所以不需要手动调用watch(),就能正常计算损失对这些变量的梯度。
with tf.GradientTape() as tape: logits = model(images, training=True) loss_value = loss_object(labels, logits) loss_history.append(loss_value.numpy().mean()) grads = tape.gradient(loss_value, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables))
内容的提问来源于stack exchange,提问作者user1245262
相关产品推荐
相关产品推荐

