tf.test.compute_gradient扰动delta相关异常行为咨询
tf.test.compute_gradient梯度校验异常问题 使用tf.test.compute_gradient校验TensorFlow对高复杂度自定义损失函数的求导正确性时,接口返回结果随扰动参数delta(下文记为dx)的取值变化,表现不符合理论预期。
数值梯度计算原理
对于一维函数f(x),tf.test.compute_gradient采用中心差分格式计算数值雅可比,计算公式为:(f(x + dx) - f(x - dx)) / (2 dx)
理论上当dx -> 0时,截断误差E(dx) = f'''(x) dx^2 / 6 -> 0,数值雅可比应收敛到解析真值,高维场景下该收敛规律基本一致。
异常现象与待解决疑问
实际测试中发现,即使是形式非常简单的函数,也会出现dx -> 0时E(dx) -> +∞的反常现象,需要明确以下问题的解决方案:
- 数值雅可比无法按理论收敛时,如何正确测试自定义函数的梯度,排查TensorFlow自动求导逻辑的错误?
- 如何区分浮点舍入误差和自动求导本身的计算错误?
delta参数的合理取值应当如何选择?delta的最优取值是否会随问题维度变化(例如存在tf.reduce_sum操作导致误差累积的场景)?
复现验证
选取正态分布对数似然函数作为测试对象,该函数三阶导数理论值为0,数值雅可比理论上应与解析值完全一致,但实际测试中max(Jth - Jnu)(解析雅可比与数值雅可比的最大绝对误差)随delta减小反而升高。
该问题在dtype=float32和dtype=float64精度、任意维度n下均可复现:当n=1、eval_jac_at=1时误差处于较低水平,但在刚性更强的场景(eval_jac_at >> 1)或高维场景(n >> 1)下误差会达到不可忽略的水平。
复现代码
import tensorflow as tf import tensorflow.math as tfm from tensorflow_probability import distributions as tfd import numpy as np import matplotlib.pyplot as plt """ 对比tf.test.compute_gradient计算的理论梯度与数值梯度随扰动delta的变化规律 测试函数为多元正态分布对数似然: log(N(loc, scale)) = \sum_i (x_i - loc_i)**2 / (2 * scale_i**2) + 常数项 """ # 特征维度 n = 1 # 数据精度 dtype = "float32" # 雅可比计算点(所有维度均取该值) eval_jac_at = 1 loc = tf.cast(0.0, dtype) # scale取值范围在0.5到1.5之间 scale = tf.cast(0.5 + np.random.rand(n), dtype) norm_dist = tfd.Normal(loc=loc, scale=scale) def f(x): """ 显式定义正态对数似然 """ return tfm.reduce_sum(-tfm.pow(x - loc, 2) / (2 * scale ** 2)) def g(x): """ 调用tfp.distributions接口计算正态对数似然 """ return tfm.reduce_sum(norm_dist.log_prob(x)) def compute_err_gradient(f, delta): """ 给定函数f和扰动delta,计算max(|数值雅可比 - 理论雅可比|) 雅可比计算点为eval_jac_at * [1, ..., 1] """ Jth, Jnu = tf.test.compute_gradient(f, [eval_jac_at * tf.ones(n, dtype)], delta) return np.max(np.abs(Jnu[0] - Jth[0])) # 生成delta序列:delta = 1 / 2**i deltas = 1 / np.power(2, np.arange(5, 14)) # 计算每个delta对应的雅可比误差 err_f = np.array([compute_err_gradient(f, delta) for delta in deltas]) err_g = np.array([compute_err_gradient(g, delta) for delta in deltas]) # 打印结果并绘图 print() print(f"{deltas=}") print(f"{err_f=}") print(f"{err_g=}") print() fig, ax = plt.subplots() ax.loglog(deltas, err_f, label="f") ax.loglog(deltas, err_g, label="g") ax.set_xlabel("delta") ax.set_ylabel("Jth - Jnu") ax.legend() ax.grid() plt.show()
测试输出
deltas=array([0.03125 , 0.015625 , 0.0078125 , 0.00390625, 0.00195312, 0.00097656, 0.00048828, 0.00024414, 0.00012207]) err_f=array([1.19209290e-07, 1.19209290e-07, 1.19209290e-07, 1.19209290e-07, 3.69548798e-06, 3.93390656e-06, 1.91926956e-05, 1.13248825e-05, 7.23600388e-05], dtype=float32) err_g=array([1.19209290e-07, 3.69548798e-06, 3.93390656e-06, 3.93390656e-06, 1.13248825e-05, 1.13248825e-05, 4.97102737e-05, 4.97102737e-05, 4.97102737e-05], dtype=float32)
内容的提问来源于stack exchange,提问作者valade aurélien
相关产品推荐
相关产品推荐

