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

如何将tf.gradients代码迁移为tf.GradientTape实现(TF2.11)

解决方案

核心问题分析

原tf.gradients(x_conv, x, x_conv)的作用是:计算x_conv对x的梯度,同时将x_conv作为**上游梯度(初始梯度)**传入,对应TensorFlow 2中GradientTape.gradient()的output_gradients参数。你得到inverse = None的原因是:x不是可训练变量,GradientTape默认不会追踪它的梯度,需要手动开启追踪。

修改后的可复现测试代码

import tensorflow as tf
import numpy as np

data = tf.random.uniform((4,3,16800), dtype=tf.float32)

with tf.GradientTape() as tape:
    x = data
    # 关键:手动让tape追踪x的梯度(因为x不是可训练变量)
    tape.watch(x)
    shape_input = x.get_shape().as_list()
    shape_fast = [np.prod(shape_input[:-1]), 1, shape_input[-1]]
    kernel_size = 1794
    paddings = [0, 0], [0, 0], [kernel_size // 2 - 1, kernel_size // 2 + 1]
    filters_kernel = tf.random.uniform((1794, 1, 16), dtype=tf.float32)
    x_reshape = tf.reshape(x, shape_fast)
    x_pad = tf.pad(x_reshape, paddings=paddings, mode='SYMMETRIC')
    x_conv = tf.nn.conv1d(x_pad, filters_kernel, stride=2,
                          padding='VALID', data_format='NCW')
# 对应原tf.gradients的三个参数:目标张量x_conv,源张量x,上游梯度x_conv
inverse = tape.gradient(x_conv, x, output_gradients=x_conv)

# 重构损失,tf.stop_gradient用法和原代码完全一致
reconstruction_loss = tf.nn.l2_loss(inverse - tf.stop_gradient(x))

对原业务代码的改写

对应你原代码中的核心两行,改写逻辑如下:

# 原代码:
# inverse = tf.gradients(x_conv, x, x_conv)[0]
# reconstruction_loss = tf.nn.l2_loss(inverse - tf.stop_gradient(x))

# 改写后(需放在GradientTape上下文外,且上下文内要watch(x))
with tf.GradientTape() as tape:
    # ... 这里是生成x_conv的前置代码
    tape.watch(x)  # 必须添加这行,确保x被追踪
    # 执行你的卷积等操作,最终得到x_conv
    # x_conv = ... 

inverse = tape.gradient(x_conv, x, output_gradients=x_conv)
reconstruction_loss = tf.nn.l2_loss(inverse - tf.stop_gradient(x))

关键说明

  • tape.watch(x):必须在GradientTape上下文内调用,告诉TensorFlow追踪x的梯度变化,否则无法计算x_conv对x的梯度。
  • output_gradients=x_conv:完全对应原tf.gradients的第三个参数,实现用x_conv作为初始梯度的反向传播计算。
  • tf.stop_gradient(x):用法和原代码一致,确保计算损失时x不参与梯度更新,只作为固定目标值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 20:40:46