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

如何在向量值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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 23:18:09