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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 12:53:28