如何在Python中实现任意函数的自动微分?mpmath相关疑问
关于自动微分与mpmath高阶导数计算的问题解答
1. 自定义Var类中children列表与链式规则的递归逻辑
自定义Var类的核心是正向模式自动微分:每个操作节点(比如加法、乘法)会记录自己的子节点(参与运算的变量/中间结果),以及该操作对每个子节点的局部导数。
以加法操作为例,[(1, self), (1, other)]的含义是:对于self + other这个节点,它对self的局部导数是1,对other的局部导数也是1。当求整体导数时,会递归遍历每个子节点:
- 先取当前节点的局部导数,再乘以子节点的导数(递归计算子节点的导数)
- 把所有子节点的贡献累加,最终得到当前节点对输入变量的导数
举个简单例子:如果x是Var实例,y = x + x,那么y的children是[(1,x), (1,x)]。求dy/dx时,就是1*(dx/dx) + 1*(dx/dx) = 1+1=2,这就是递归链式规则的实际执行过程。
2. 自定义Var类的局限性:为何无法支持任意函数与高阶导
你说它像“伪装的符号计算”其实不准确——它本质是计算图追踪的自动微分,但问题在于你的自定义类只实现了部分操作(比如加减乘除),没有覆盖所有需要的数学函数(比如cos)。
要支持任意函数和高阶导,你需要给每个用到的数学函数都实现Var类的包装:
- 比如要支持
cos(x),你需要定义一个接受Var实例的cos函数,返回新的Var节点,其children为[(-sin(x.value), x)](因为cos(x)对x的导数是-sin(x)) - 高阶导的本质是对“导数结果”再做一次自动微分,如果一阶导的计算过程中用到了未实现的函数(比如二阶导需要计算cos的导数sin),而你的Var类没实现sin函数,就会报错。
这种自定义类的方式扩展性差,必须手动实现所有需要的操作,没法直接支持任意函数。
3. Python中实现任意函数自动微分的方案,以及mpmath的替代/结合方式
无需自定义类:用成熟的自动微分库
不用自己造轮子,Python有很多现成的自动微分库可以直接支持任意纯Python函数的高阶导数计算,比如:
- JAX:支持正向/反向模式自动微分,高阶导数只需嵌套调用
jax.grad即可,示例:import jax import jax.numpy as jnp def f(x): return jnp.cos(x) + x**2 # 一阶导数 df_dx = jax.grad(f) # 二阶导数 d2f_dx2 = jax.grad(df_dx) # 计算x=1时的二阶导数 print(d2f_dx2(1.0)) - PyTorch:同样支持自动微分,适合深度学习场景,但也能用于普通数学函数的高阶导计算。
mpmath的相关方案
mpmath本身主打高精度数值计算,原生没有自动微分功能:
- 它的
mpmath.diff函数是数值微分(用有限差分法),精度会随着阶数升高而下降,不适合高阶导数的精确计算。 - 如果一定要结合mpmath,可以用自动微分库包装mpmath的函数,但需要确保库支持(比如JAX对部分mpmath函数兼容有限,更推荐用JAX自带的高精度数值函数)。
内容的提问来源于stack exchange,提问作者ShoutOutAndCalculate
相关产品推荐
相关产品推荐

