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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 00:27:27