手动实现链式法则(Chain Rule)结果异常,求排查问题原因
链式法则实现错误排查
你尝试用纯Python实现链式法则,但计算结果和基准值不符,你的代码如下:
def calc_derivative(func, x, dx): return (func(x + dx) - func(x)) / dx
def chain_diff(func1, func2, x, dx): return calc_derivative(func1(func2), x, dx) * calc_derivative(func2, x, dx)
def func1(f): g = np.log def apply(x): return g(f(x)) return apply def func2(x): return x ** 4
span = 100 y, dy = [], [] tst = [] for x in range(1, span): x *= .1 dx = np.sqrt(2e-15) * x dy.append(chain_diff(func1, func2, x, dx)) y.append(func1(func2)(x)) tst.append(4 / x) plt.plot(dy) plt.plot(tst) plt.plot(y) plt.ylim((-10, 10)) plt.legend(("$dy$", "tst plot","y")) plt.show()
错误原因
你核心的问题出在chain_diff函数的逻辑上,完全搞错了链式法则的应用方式:
链式法则的数学定义是 $\frac{d}{dx}[f(g(x))] = f'(g(x)) \times g'(x)$,而你的代码里:
calc_derivative(func1(func2), x, dx)已经直接算出了复合函数$f(g(x))$在$x$处的导数(这本身就是链式法则的结果)- 你又多乘了一次
calc_derivative(func2, x, dx),相当于把正确结果再乘了一遍$g'(x)$,最终结果变成$\frac{d}{dx}[f(g(x))] \times g'(x)$,自然和基准值$\frac{4}{x}$完全对不上。
另外,func1的嵌套写法完全没必要,反而把外层函数(log)的逻辑藏起来了,增加了理解难度。
修正后的代码
核心函数修正
import numpy as np import matplotlib.pyplot as plt def calc_derivative(func, x, dx): return (func(x + dx) - func(x)) / dx # 正确实现链式法则:外层函数在g(x)处的导数 × 内层函数在x处的导数 def chain_diff(outer_func, inner_func, x, dx): # 先算内层函数在x处的值:g(x) inner_val = inner_func(x) # 算外层函数在g(x)点的导数:f'(g(x)) outer_deriv = calc_derivative(outer_func, inner_val, dx) # 算内层函数在x点的导数:g'(x) inner_deriv = calc_derivative(inner_func, x, dx) # 链式法则相乘 return outer_deriv * inner_deriv # 内层函数保持不变 def func2(x): return x ** 4
验证代码
span = 100 y, dy = [], [] tst = [] for x in range(1, span): x *= .1 dx = np.sqrt(2e-15) * x # 直接传入np.log作为外层函数,func2作为内层函数 dy.append(chain_diff(np.log, func2, x, dx)) y.append(np.log(func2(x))) tst.append(4 / x) plt.plot(dy, label="$dy$") plt.plot(tst, label="基准导数") plt.plot(y, label="$y$") plt.ylim((-10, 10)) plt.legend() plt.show()
结果说明
修正后,dy的曲线会和基准值$\frac{4}{x}$几乎重合——因为$y = \log(x4)$的导数就是$\frac{4x3}{x^4} = \frac{4}{x}$,完全符合预期。
内容的提问来源于stack exchange,提问作者Nabla
相关产品推荐
相关产品推荐

