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

如何将jax.custom_vjp用于含非JAX类型(如SymPy表达式)输入的函数?

问题

尝试用JAX的jax.custom_vjp给接收SymPy表达式的函数定义自定义梯度时,遇到错误:JAX不支持非JAX类型作为grad、jit或custom_vjp转换函数的输入。以下是最小复现代码:

import jax
import sympy as sm

# Define symbols and expression
x, y, z = sm.symbols('x y z')
expr = x**2 + 2*y + z

# Attempt to iterate over expr (this will cause an error)
try:
    for term in expr:
        print(term)
except TypeError as e:
    print(f"Error: {e}")

# Define a function that takes a SymPy expression and a value
def sympy_function(expr, x_value):
    x = sm.Symbol('x')
    result = expr.subs(x, x_value)
    return float(result)

# Attempt to apply custom_vjp
sympy_function = jax.custom_vjp(sympy_function)

def sympy_function_fwd(expr, x_value):
    y = sympy_function(expr, x_value)
    return y, (expr, x_value)

def sympy_function_bwd(residual, grad_y):
    expr, x_value = residual
    x = sm.Symbol('x')
    derivative_expr = sm.diff(expr, x)
    grad_x_value = float(derivative_expr.subs(x, x_value))
    grad_expr = None
    return grad_expr, grad_y * grad_x_value

sympy_function.defvjp(sympy_function_fwd, sympy_function_bwd)

# Test the function
x = sm.Symbol('x')
expr = x**2 + 3*x + 2
x_value = 1.0

# This will raise an error
y = sympy_function(expr, x_value)

运行后报错:

TypeError: Value x**2 + 3*x + 2 with type <class 'sympy.core.add.Add'> is not a valid JAX type

如何将jax.custom_vjp用于含SymPy表达式这类非JAX类型输入的函数?是否有方法规避该限制?

解决方案

方法1:将SymPy表达式标记为静态参数

JAX允许将非JAX类型标记为静态参数,这类参数不会被JAX追踪,也不参与自动微分。可以结合jax.jit的static_argnums选项与jax.custom_vjp实现需求:

import jax
import sympy as sm

def sympy_function(expr, x_value):
    x = sm.Symbol('x')
    result = expr.subs(x, x_value)
    return float(result)

@jax.custom_vjp
def sympy_function_static(expr, x_value):
    return sympy_function(expr, x_value)

def fwd(expr, x_value):
    y = sympy_function(expr, x_value)
    return y, (expr, x_value)

def bwd(residual, grad_y):
    expr, x_value = residual
    x = sm.Symbol('x')
    derivative_expr = sm.diff(expr, x)
    grad_x = float(derivative_expr.subs(x, x_value))
    return None, grad_y * grad_x  # 静态参数无梯度,返回None

sympy_function_static.defvjp(fwd, bwd)

# 标记第0个参数(expr)为静态参数
sympy_function_jitted = jax.jit(sympy_function_static, static_argnums=(0,))

# 测试
x = sm.Symbol('x')
expr = x**2 + 3*x + 2
x_value = 1.0

y = sympy_function_jitted(expr, x_value)
print(f"函数输出: {y}")

# 计算x_value的梯度
grad_fn = jax.grad(sympy_function_jitted, argnums=1)
grad_result = grad_fn(expr, x_value)
print(f"x_value的梯度: {grad_result}")

方法2:提前将SymPy表达式编译为JAX函数

把SymPy表达式转换成JAX可追踪的函数,避免直接传递SymPy对象给JAX转换后的函数:

import jax
import jax.numpy as jnp
import sympy as sm

# 将SymPy表达式转为JAX函数
def sympy_to_jax(expr, var):
    return sm.lambdify(var, expr, modules='jax')

@jax.custom_vjp
def jax_compatible_function(expr, x_value):
    x = sm.Symbol('x')
    jax_fn = sympy_to_jax(expr, x)
    return jax_fn(x_value)

def fwd(expr, x_value):
    x = sm.Symbol('x')
    jax_fn = sympy_to_jax(expr, x)
    y = jax_fn(x_value)
    # 提前计算导数并转为JAX函数
    deriv_expr = sm.diff(expr, x)
    jax_deriv_fn = sympy_to_jax(deriv_expr, x)
    return y, (jax_deriv_fn, x_value)

def bwd(residual, grad_y):
    jax_deriv_fn, x_value = residual
    grad_x = jax_deriv_fn(x_value)
    return None, grad_y * grad_x  # expr为静态参数,梯度返回None

jax_compatible_function.defvjp(fwd, bwd)

# 测试
x = sm.Symbol('x')
expr = x**2 + 3*x + 2
x_value = jnp.array(1.0)  # 使用JAX数组类型

y = jax_compatible_function(expr, x_value)
print(f"函数输出: {y}")

# 计算x_value的梯度
grad_fn = jax.grad(jax_compatible_function, argnums=1)
grad_result = grad_fn(expr, x_value)
print(f"x_value的梯度: {grad_result}")

关键说明

  • JAX仅追踪原生类型(如JAX数组、标量),非JAX类型必须标记为静态参数,否则会触发类型错误。
  • 静态参数在JAX转换过程中被视为常量,无法对其求导,因此梯度需返回None。
  • 若需对与SymPy表达式相关的参数求导,建议提前将表达式转为JAX函数,仅将可变参数(如x_value)作为JAX类型输入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 02:25:12