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

tf.GradientTape梯度调用次数是否受限?四层嵌套出现NaN异常

核心结论

tf.GradientTape没有嵌套层数的机制上限,只要内存充足,嵌套任意层数做高阶微分都是官方支持的。你遇到的边角位置NaN问题和层数限制无关,是代码写法带来的计算路径释放+数值稳定性问题。

具体原因
  • 非持久磁带的资源提前释放
    你当前代码中只有两层嵌套磁带设置了persistent=True,另外两层是默认非持久模式。非持久磁带在退出with作用域的瞬间就会释放持有的计算图资源,四层嵌套下,计算最外层梯度时需要访问内层非持久磁带保存的中间节点路径,这部分资源已经被释放,就会产生无梯度的NaN值。三层嵌套时不会触发是因为梯度计算路径更短,磁带释放前所有需要的梯度已经计算完成。
    之所以只有[0,0]和[-1,-1]两个位置出问题,是因为你用tf.repeat+转置生成的计算网格中,这两个点刚好是网格对角端点,计算路径最先被回收,其他位置的路径释放晚,刚好赶上梯度计算完成。
  • 对角奇点的数值不稳定
    你在协方差函数中手动实现的softplus转换tf.math.log(1+tf.math.exp(x))存在数值上溢风险,同时在x=x_、t=t_的对角位置,指数项输入为0,四阶微分计算时会碰到0/0型的未定义奇点,端点位置没有相邻元素的数值平滑,直接就会算出NaN。
可直接运行的修复代码

只需要修改两处:一是所有嵌套磁带统一设为持久模式,梯度算完手动删除释放内存;二是替换手动实现的softplus为内置稳定版本,给距离项加极小值扰动避开奇点。

import tensorflow as tf
import matplotlib.pyplot as plt
import numpy as np
import math
from scipy import io 

tf.random.set_seed(1234)
np.random.seed(1234)

n_u=20
n_f=30

t_u=tf.cast(tf.linspace(0,1,n_u),tf.float32)
t_u=tf.reshape(t_u,shape=(n_u,1))
t_u=tf.Variable(t_u)

x_u=tf.cast(tf.linspace(0,1,n_u),tf.float32)
x_u=tf.reshape(x_u,shape=(n_u,1))
x_u=tf.Variable(x_u)

t_f=tf.cast(tf.linspace(0,1,n_f),tf.float32)
t_f=tf.reshape(t_f,shape=(n_f,1))
t_f=tf.Variable(t_f)

x_f=tf.cast(tf.linspace(0,1,n_f),tf.float32)
x_f=tf.reshape(x_f,shape=(n_f,1))
x_f=tf.Variable(x_f)

def cov(x,x_,t,t_,sigma=0.5, weight_x=-1, weight_t=-3,repeat = False):
    # 用内置tf.nn.softplus替换手动实现,避免指数上溢
    sigma = tf.nn.softplus(tf.constant(sigma,dtype=tf.float32))
    weight_x = tf.nn.softplus(tf.constant(weight_x,dtype=tf.float32))
    weight_t = tf.nn.softplus(tf.constant(weight_t,dtype=tf.float32))

    if repeat:
        x=tf.repeat(x,x_.shape[0],axis=1)
        x_=tf.transpose(tf.repeat(x_,x.shape[0],axis=1))
        t=tf.repeat(t,t_.shape[0],axis=1)
        t_=tf.transpose(tf.repeat(t_,t.shape[0],axis=1))
    # 距离项加1e-8极小扰动,避开严格0值带来的四阶导奇点
    dist_x = (x - x_)**2 + 1e-8
    dist_t = (t - t_)**2 + 1e-8
    return sigma**2 * tf.math.exp(-0.5 * weight_x * dist_x - 0.5 * weight_t * dist_t)

def K_12(x,x_,t,t_):
    x=tf.repeat(x,x_.shape[0],axis=1)
    x_=tf.transpose(tf.repeat(x_,x.shape[0],axis=1))
    t=tf.repeat(t,t_.shape[0],axis=1)
    t_=tf.transpose(tf.repeat(t_,t.shape[0],axis=1))

    # 所有嵌套磁带统一设为持久模式,避免中间路径提前释放
    with tf.GradientTape(persistent=True) as g:
        g.watch(x_)
        with tf.GradientTape(persistent=True) as gg:
            gg.watch(x_)
            gg.watch(t_)
            with tf.GradientTape(persistent=True) as ggg:
                ggg.watch(x)
                with tf.GradientTape(persistent=True) as gggg:
                    gggg.watch(x)
                    gggg.watch(t)
                    K_uu=cov(x,x_,t,t_)
                K_x=gggg.gradient(K_uu,x)
                K_t=gggg.gradient(K_uu,t)
            K_xx=ggg.gradient(K_x,x)
            K_f=K_t - K_xx
        K_f_x=gg.gradient(K_f,x_)
        K_f_t=gg.gradient(K_f,t_)
    K_f_xx=g.gradient(K_f_x,x_)
    # 手动删除持久磁带释放内存
    del g, gg, ggg, gggg
    return K_t - K_f_xx

print(K_12(x_u,x_f,t_u,t_f))

修改后运行不会再出现边角位置的NaN值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 00:42:22