如何将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
相关产品推荐
相关产品推荐

