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
相关产品推荐
相关产品推荐

