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

JAX与mpmath兼容性问题:自动微分调用函数遇类型错误

解决JAX与mpmath兼容性导致的自动微分错误

问题根源

mpmath是基于Python标量的数值/符号计算库,完全不兼容JAX的自动微分追踪机制:

  • JAX在微分时会将输入标记为抽象追踪数组,而mpmath函数仅能处理Python原生标量
  • complex(x)操作试图将JAX抽象数组转换为Python复数,直接触发Concretization type error——JAX禁止在追踪期间将抽象数组"具体化"为原生Python类型

方案1:替换为JAX原生兼容的函数(推荐)

JAX的jax.scipy.special模块已经实现了多对数函数polylog,完全支持自动微分,直接替换即可解决问题:

import jax
import jax.scipy.special as jsp

def func(x):
    return jsp.polylog(2, x)  # 无需手动转complex,JAX会自动处理复数运算

jac = jax.jacobian(func)
# 测试调用
print(jac(0.5))  # 输出对应导数

方案2:用jax.pure_callback+自定义VJP包装mpmath函数(仅当JAX无对应实现时使用)

如果必须依赖mpmath的特定实现,可通过自定义VJP(向量雅可比积)将其包装为JAX可追踪的操作,同时手动定义梯度函数(因为mpmath本身不提供JAX兼容的梯度):

多对数函数$\text{Li}_2(z)$的导数公式为:$\frac{d}{dz}\text{Li}_2(z) = \frac{\ln(1-z)}{z}$,基于此我们可以手动实现梯度逻辑:

import jax
from mpmath import polylog as polylog_mpmath, log

def func_mpmath(x):
    # 处理标量输入(自定义VJP会自动将JAX数组转为标量传入)
    return polylog_mpmath(2, complex(x))

def grad_func_mpmath(x):
    # 实现Li2的导数计算
    return log(1 - complex(x)) / complex(x)

# 包装为JAX可微分函数
func_jax = jax.custom_vjp(func_mpmath)

# 定义前向和反向传播规则
@func_jax.defvjp
def func_vjp(x):
    y = func_mpmath(x)
    def vjp(grad_output):
        return grad_output * grad_func_mpmath(x),
    return y, vjp

# 计算雅可比矩阵
jac = jax.jacobian(func_jax)
# 测试调用
print(jac(0.5))

说明

  • jax.custom_vjp用于自定义函数的正向和反向传播逻辑
  • 前向传播调用mpmath的原函数,反向传播使用手动实现的导数公式
  • 该方法仅适用于标量或简单输入场景,复杂输入下需要额外处理数组遍历

内容的提问来源于stack exchange,提问作者Fr6

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 05:38:27