如何在向量值ODE中使用与解释JAX VJP
JAX中ODE反向模式VJP的使用困惑及解决方案
问题背景
我正在学习用JAX求解向量值ODE函数的Jacobian,尝试用官方jax.experimental.ode.odeint和diffrax库复现自定义积分器的正反向Jacobian功能,但二者默认采用反向模式Vector-Jacobian Product(VJP),而非教程中的正向模式Jacobian Vector Product(JVP)。我对VJP的概念及向量值ODE场景下的输入形式存在误解,调用VJP时出现错误,同时想明确VJP的作用和正确用法。
出错代码片段
import matplotlib.pyplot as plt from jax.config import config config.update("jax_enable_x64", True) import jax.numpy as jnp from jax import jit, jvp, vjp from jax.experimental.ode import odeint from diffrax import diffeqsolve, ODETerm, PIDController, SaveAt, Dopri5, NoAdjoint # 定义ODE右端函数(向量值) def f(state, t, args): x, y, z = state rho, sigma, beta = args return jnp.array([sigma * (y - x), x * (rho - z) - y, x * y - beta * z]) # 封装ODE积分过程(待求Jacobian的函数) def evolve(y0, rho, sigma, beta): return odeint(f, y0, tarr, (rho, sigma, beta)) # 初始化条件与参数 y0 = jnp.array([5., 5., 5.]) tarr = jnp.linspace(0, 1., 1000) rho = 28. sigma = 10. beta = 8/3. # 验证积分函数正常工作 ys = evolve(y0, rho, sigma, beta) fig, ax = plt.subplots(1,figsize=(6,4),dpi=150,subplot_kw={'projection':'3d'}) ax.plot(ys.T[0],ys.T[1],ys.T[2],'b-',lw=0.5) # 尝试计算反向模式VJP vjp_ys, vjp_evolve = vjp(evolve,y0,rho,sigma,beta) print(jnp.array_equal(ys,vjp_ys)) # 定义输入扰动 delta_y0 = jnp.array([0., 0., 0.]) delta_rho = 0. delta_sigma = 0. delta_beta = 1. # 此处调用出错 vjp_evolve(delta_y0,delta_rho,delta_sigma,delta_beta)
错误信息
TypeError: The function returned by `jax.vjp` applied to evolve was called with 4 arguments, but functions returned by `jax.vjp` must be called with a single argument corresponding to the single value returned by evolve (even if that returned value is a tuple or other container). For example, if we have: def f(x): return (x, x) _, f_vjp = jax.vjp(f, 1.0) the function `f` returns a single tuple as output, and so we call `f_vjp` with a single tuple as its argument: x_bar, = f_vjp((2.0, 2.0)) If we instead call `f_vjp(2.0, 2.0)`, with the values 'splatted out' as arguments rather than in a tuple, this error can arise.
已实现的正向模式JVP代码(Diffrax)
# Diffrax兼容的ODE右端函数(时间在前) def f_diffrax(t, state, args): x, y, z = state rho, sigma, beta = args return jnp.array([sigma * (y - x), x * (rho - z) - y, x * y - beta * z]) # 封装Diffrax积分过程 terms = ODETerm(f_diffrax) t0 = 0.0 t1 = 1.0 dt0 = None max_steps = 16**3 tsave = SaveAt(ts=tarr,dense=True) def evolve_diffrax(y0, rho, sigma, beta): return diffeqsolve(terms,Dopri5(),t0,t1,dt0,y0,jnp.array([rho,sigma,beta]),saveat=tsave, stepsize_controller=PIDController(rtol=1.4e-8,atol=1.4e-8),max_steps=max_steps,adjoint=NoAdjoint()) # 计算正向JVP diffrax_ys, diffrax_delta_ys = jvp(evolve_diffrax, (y0,rho,sigma,beta),(delta_y0,delta_rho,delta_sigma,delta_beta)) # 提取解数组 diffrax_ys = diffrax_ys.ys diffrax_delta_ys = diffrax_delta_ys.ys # 可视化 fig, ax = plt.subplots(1,figsize=(6,4),dpi=150,subplot_kw={'projection':'3d'}) ax.plot(diffrax_ys.T[0],diffrax_ys.T[1],diffrax_ys.T[2],color='violet',lw=0.5) ax.quiver(diffrax_ys.T[0][::10],diffrax_ys.T[1][::10],diffrax_ys.T[2][::10], diffrax_delta_ys.T[0][::10],diffrax_delta_ys.T[1][::10],diffrax_delta_ys.T[2][::10])
解决方案与概念解析
1. VJP与JVP的核心差异
- JVP(正向模式):给定输入的扰动(Δ输入),计算输出的扰动(Δ输出 = J·Δ输入),适合输入维度远小于输出维度的场景。你用Diffrax的
NoAdjoint实现的就是这种模式,对应ODE解对初始条件/参数的正向敏感性。 - VJP(反向模式):给定输出的“伴随”(即输出的梯度/扰动,Δ输出),计算输入的伴随(Δ输入 = Jᵀ·Δ输出),适合输出维度远小于输入维度的场景,常用于损失函数对初始条件/参数的梯度反向传播。
2. 错误根源:VJP函数调用方式错误
你调用vjp_evolve时传入了多个输入扰动参数,但VJP返回的函数仅接受一个参数——该参数必须是与原函数输出同形状的张量,代表输出的扰动,而非输入的扰动。
原函数evolve的输出是形状为(1000, 3)的ODE解序列,因此需要传入同形状的扰动张量来计算输入(y0, rho, sigma, beta)对应的伴随值。
3. 正确调用VJP的示例代码
修改VJP调用部分如下:
# 定义输出扰动:例如仅对最后一个时间步的x分量施加单位扰动 delta_ys = jnp.zeros_like(ys) delta_ys = delta_ys.at[-1].set(jnp.array([1.0, 0.0, 0.0])) # 调用VJP函数,传入输出扰动,得到输入的伴随值 delta_y0_vjp, delta_rho_vjp, delta_sigma_vjp, delta_beta_vjp = vjp_evolve(delta_ys) print("y0的伴随值:", delta_y0_vjp) print("rho的伴随值:", delta_rho_vjp) print("sigma的伴随值:", delta_sigma_vjp) print("beta的伴随值:", delta_beta_vjp)
这段代码的含义是:如果希望ODE解在最后一步的x分量增加1,初始条件和各个参数需要产生多大的变化(本质是Jacobian转置与输出扰动的乘积)。
4. Diffrax中VJP的正确使用方式
Diffrax默认支持反向模式自动微分,只需将积分函数的输出转为数值数组(而非Solution对象),再用jax.vjp包裹即可:
def evolve_diffrax_output(y0, rho, sigma, beta): sol = diffeqsolve(terms,Dopri5(),t0,t1,dt0,y0,jnp.array([rho,sigma,beta]),saveat=tsave, stepsize_controller=PIDController(rtol=1.4e-8,atol=1.4e-8),max_steps=max_steps) return sol.ys # 直接返回解数组 # 计算VJP vjp_ys_diffrax, vjp_evolve_diffrax = vjp(evolve_diffrax_output, y0, rho, sigma, beta) # 定义输出扰动 delta_ys_diffrax = jnp.zeros_like(vjp_ys_diffrax) delta_ys_diffrax = delta_ys_diffrax.at[-1].set(jnp.array([1.0, 0.0, 0.0])) # 获取输入伴随值 delta_y0_dx, delta_rho_dx, delta_sigma_dx, delta_beta_dx = vjp_evolve_diffrax(delta_ys_diffrax)
总结
- JVP是“输入扰动→输出扰动”,适合求解ODE解对初始条件/参数的正向敏感性;
- VJP是“输出扰动→输入伴随”,适合求解损失函数对初始条件/参数的梯度;
- 调用VJP返回的函数时,必须传入与原函数输出同形状的单个张量,而非多个输入扰动参数。
内容的提问来源于stack exchange,提问作者Jim Raynor
相关产品推荐
相关产品推荐

