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

JAX中非单位向量的向量-雅可比乘积(VJP)含义与验证疑问

反向模式自动微分中向量-雅可比乘积(VJP)的含义与验证误区

问题描述

我对反向模式自动微分中,向量值函数使用非单位行向量计算向量-雅可比乘积(VJP)的含义存在困惑。已知用单位向量可以提取雅可比矩阵的行(对应单个输出分量对输入的梯度),但当输入非单位向量如[2,0]时,我原以为得到的是能让输出第一分量翻倍的参数扰动,却无法通过直接扰动参数验证结果,想搞清楚操作中的错误。

代码示例

from jax.config import config
config.update("jax_enable_x64", True)
import jax.numpy as jnp
from jax import vjp, jacrev

# 定义向量值函数(3输入→2输出)
def vector_func(args):
    x,y,z = args
    a = 2*x**2 + 3*y**2 + 4*z**2
    b = 4*x*y*z
    return jnp.array([a, b])

# 定义输入
x = 2.0
y = 3.0
z = 4.0

# 在基准点计算向量-雅可比乘积
val, func_vjp = vjp(vector_func, (x, y, z))
print(val) 
# [99,96]

# 用单位向量提取输出分量的梯度
v1 = jnp.array([1.0, 0.0])  # 提取第一分量对输入的梯度(雅可比第一行)
v2 = jnp.array([0.0, 1.0])  # 提取第二分量对输入的梯度(雅可比第二行)

gradient1 = func_vjp(v1)
print(gradient1)
# [8, 18, 32]

gradient2 = func_vjp(v2)
print(gradient2)
# [48,32,24]

# 尝试用非单位向量[2,0]获取参数扰动
print(func_vjp(jnp.array([2.0,0.0])))
# [16,36,64] 

# 验证尝试均失败
print(vector_func([16,36,64]))
# [20784, 147456]

print(vector_func([x*16,y*36,z*64]))
# [299184., 3538944.]

误区解析

你对VJP的含义理解出现了核心偏差:

  • VJP的本质是行向量与雅可比矩阵的乘积,即v · J,其中v是输出侧的线性组合系数,J是函数的雅可比矩阵(行对应输出分量,列对应输入分量)。
  • 当输入v=[2,0]时,结果是2 * J[0,:]——也就是第一输出分量梯度的2倍,这是输入侧梯度的线性组合,而非参数的扰动值或缩放因子。
  • 你之前的验证逻辑错误在于:把VJP结果直接当作新的参数输入,或者用参数乘以VJP结果,这完全不符合VJP的数学定义。

正确验证方法

方法1:直接对比手动计算的向量-雅可比乘积

通过jacrev计算完整雅可比矩阵,再手动计算行向量与雅可比的乘积,验证是否与VJP结果一致:

# 计算完整雅可比矩阵
jac_matrix = jacrev(vector_func)(x, y, z)
print("雅可比矩阵:")
print(jac_matrix)
# [[8. 18. 32.]
#  [48. 32. 24.]]

# 手动计算v·J
v = jnp.array([2.0, 0.0])
manual_vjp_result = jnp.dot(v, jac_matrix)
print("\n手动计算VJP结果:")
print(manual_vjp_result)  # [16. 36. 64.]

# 对比VJP输出
vjp_result = func_vjp(v)[0]
print("\nJAX VJP结果:")
print(vjp_result)  # [16. 36. 64.]

两者结果完全一致,说明VJP的计算是正确的。

方法2:一阶泰勒近似验证方向导数

VJP的结果g = v·J对应的是:输出侧线性组合v·f(x)的梯度为g。可以通过一阶泰勒近似验证:对于微小扰动ε,有v·f(x + εu) ≈ v·f(x) + ε g·u(其中u是任意输入方向)。

取u为单位向量,ε=1e-3验证:

v = jnp.array([2.0, 0.0])
g = func_vjp(v)[0]
epsilon = 1e-3
u = jnp.array([1.0, 0.0, 0.0])  # 沿x轴方向扰动

# 计算左侧:v·f(x + εu)
perturbed_input = (x + epsilon*u[0], y + epsilon*u[1], z + epsilon*u[2])
left = jnp.dot(v, vector_func(perturbed_input))

# 计算右侧:v·f(x) + ε*g·u
right = jnp.dot(v, val) + epsilon * jnp.dot(g, u)

print(f"左侧值:{left}")
print(f"右侧近似值:{right}")
# 输出近似相等,验证成立

内容的提问来源于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.26 09:03:14